Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
llguidance.cpp271 linesDownload Raw Back to common
1#include "sampling.h"2#include "log.h"3 4#ifdef LLAMA_USE_LLGUIDANCE5 6#    include "llguidance.h"7#    include <cmath>8 9struct llama_sampler_llg {10    const llama_vocab * vocab;11    std::string         grammar_kind;12    std::string         grammar_data;13    LlgTokenizer *      tokenizer;14    LlgConstraint *     grammar;15    LlgMaskResult       llg_res;16    bool                has_llg_res;17};18 19static LlgConstraint * llama_sampler_llg_new(LlgTokenizer * tokenizer, const char * grammar_kind,20                                             const char * grammar_data) {21    LlgConstraintInit cinit;22    llg_constraint_init_set_defaults(&cinit, tokenizer);23    const char * log_level = getenv("LLGUIDANCE_LOG_LEVEL");24    if (log_level && *log_level) {25        cinit.log_stderr_level = atoi(log_level);26    }27    auto c = llg_new_constraint_any(&cinit, grammar_kind, grammar_data);28    if (llg_get_error(c)) {29        LOG_ERR("llg error: %s\n", llg_get_error(c));30        llg_free_constraint(c);31        return nullptr;32    }33    return c;34}35 36static const char * llama_sampler_llg_name(const llama_sampler * /*smpl*/) {37    return "llguidance";38}39 40static void llama_sampler_llg_accept_impl(llama_sampler * smpl, llama_token token) {41    auto * ctx = (llama_sampler_llg *) smpl->ctx;42    if (ctx->grammar) {43        LlgCommitResult res;44        llg_commit_token(ctx->grammar, token, &res);45        ctx->has_llg_res = false;46    }47}48 49static void llama_sampler_llg_apply(llama_sampler * smpl, llama_token_data_array * cur_p) {50    auto * ctx = (llama_sampler_llg *) smpl->ctx;51    if (ctx->grammar) {52        if (!ctx->has_llg_res) {53            if (llg_compute_mask(ctx->grammar, &ctx->llg_res) == 0) {54                ctx->has_llg_res = true;55            } else {56                LOG_ERR("llg error: %s\n", llg_get_error(ctx->grammar));57                llg_free_constraint(ctx->grammar);58                ctx->grammar = nullptr;59            }60        }61        if (ctx->has_llg_res) {62            if (ctx->llg_res.is_stop) {63                for (size_t i = 0; i < cur_p->size; ++i) {64                    if (!llama_vocab_is_eog(ctx->vocab, cur_p->data[i].id)) {65                        cur_p->data[i].logit = -INFINITY;66                    }67                }68            } else {69                const uint32_t * mask = ctx->llg_res.sample_mask;70                for (size_t i = 0; i < cur_p->size; ++i) {71                    auto token = cur_p->data[i].id;72                    if ((mask[token / 32] & (1 << (token % 32))) == 0) {73                        cur_p->data[i].logit = -INFINITY;74                    }75                }76            }77        }78    }79}80 81static void llama_sampler_llg_reset(llama_sampler * smpl) {82    auto * ctx = (llama_sampler_llg *) smpl->ctx;83    if (!ctx->grammar) {84        return;85    }86 87    auto * grammar_new = llama_sampler_llg_new(ctx->tokenizer, ctx->grammar_kind.c_str(), ctx->grammar_data.c_str());88    llg_free_constraint(ctx->grammar);89    ctx->grammar     = grammar_new;90    ctx->has_llg_res = false;91}92 93static llama_sampler * llama_sampler_llg_clone(const llama_sampler * smpl) {94    const auto * ctx = (const llama_sampler_llg *) smpl->ctx;95 96    auto * result = llama_sampler_init_llg(ctx->vocab, nullptr, nullptr);97 98    // copy the state99    {100        auto * result_ctx = (llama_sampler_llg *) result->ctx;101 102        if (ctx->grammar) {103            result_ctx->grammar_kind = ctx->grammar_kind;104            result_ctx->grammar_data = ctx->grammar_data;105            result_ctx->grammar      = llg_clone_constraint(ctx->grammar);106            result_ctx->tokenizer    = llg_clone_tokenizer(ctx->tokenizer);107        }108    }109 110    return result;111}112 113static void llama_sampler_llg_free(llama_sampler * smpl) {114    const auto * ctx = (llama_sampler_llg *) smpl->ctx;115 116    if (ctx->grammar) {117        llg_free_constraint(ctx->grammar);118        llg_free_tokenizer(ctx->tokenizer);119    }120 121    delete ctx;122}123 124static llama_sampler_i llama_sampler_llg_i = {125    /* .name   = */ llama_sampler_llg_name,126    /* .accept = */ llama_sampler_llg_accept_impl,127    /* .apply  = */ llama_sampler_llg_apply,128    /* .reset  = */ llama_sampler_llg_reset,129    /* .clone  = */ llama_sampler_llg_clone,130    /* .free   = */ llama_sampler_llg_free,131};132 133static size_t llama_sampler_llg_tokenize_fn(const void * user_data, const uint8_t * bytes, size_t bytes_len,134                                            uint32_t * output_tokens, size_t output_tokens_len) {135    const llama_vocab * vocab = (const llama_vocab *) user_data;136    int                 r     = 0;137    try {138        r = llama_tokenize(vocab, (const char *) bytes, bytes_len, (int32_t *) output_tokens, output_tokens_len, false,139                           true);140    } catch (const std::exception & e) {141        GGML_ABORT("llama_tokenize failed: %s\n", e.what());142    }143    if (r < 0) {144        return -r;145    }146    return r;147}148 149static LlgTokenizer * llama_sampler_llg_new_tokenizer(const llama_vocab * vocab) {150    // TODO store the tokenizer in the vocab somehow151    static const llama_vocab * vocab_cache;152    static LlgTokenizer *      tokenizer_cache;153 154    if (vocab_cache == vocab) {155        return llg_clone_tokenizer(tokenizer_cache);156    }157 158    auto tok_eos = llama_vocab_eot(vocab);159    if (tok_eos == LLAMA_TOKEN_NULL) {160        tok_eos = llama_vocab_eos(vocab);161    }162 163    size_t vocab_size = llama_vocab_n_tokens(vocab);164 165    auto token_lens       = new uint32_t[vocab_size];166    // we typically have ~7 bytes per token; let's go on the safe side here167    auto token_bytes_size = vocab_size * 16 + 1024 * 1024;168    auto token_bytes      = new uint8_t[token_bytes_size];169 170    size_t offset = 0;171    for (size_t i = 0; i < vocab_size; i++) {172        size_t max_token = 1024;173        if (token_bytes_size - offset < max_token) {174            GGML_ABORT("token_bytes buffer too small\n");175        }176 177        llama_token token = i;178        auto        dp    = (char *) token_bytes + offset;179        auto        size  = llama_detokenize(vocab, &token, 1, dp, max_token, false, false);180        if (size < 0) {181            GGML_ABORT("llama_detokenize failed\n");182        }183        if (size == 0) {184            size = llama_detokenize(vocab, &token, 1, dp + 1, max_token - 1, false, true);185            if (size < 0) {186                GGML_ABORT("llama_detokenize failed\n");187            }188            if (size != 0) {189                *dp = '\xff';  // special token prefix marker190                size += 1;191            }192        }193 194        token_lens[i] = size;195        offset += size;196    }197 198    LlgTokenizerInit tinit = {199        /* .vocab_size                         = */ (uint32_t) vocab_size,200        /* .tok_eos                            = */ (uint32_t) tok_eos,201        /* .token_lens                         = */ token_lens,202        /* .token_bytes                        = */ token_bytes,203        /* .tokenizer_json                     = */ nullptr,204        /* .tokenize_assumes_string            = */ true,205        /* .tokenize_fn                        = */ llama_sampler_llg_tokenize_fn,206        /* .use_approximate_greedy_tokenize_fn = */ false,207        /* .tokenize_user_data                 = */ vocab,208    };209 210    char           error_buffer[1024];211    LlgTokenizer * tokenizer = llg_new_tokenizer(&tinit, error_buffer, sizeof(error_buffer));212 213    delete[] token_bytes;214    delete[] token_lens;215 216    if (tokenizer == nullptr) {217        LOG_ERR("llg tokenizer error: %s\n", error_buffer);218        return tokenizer;219    }220 221    if (tokenizer_cache) {222        llg_free_tokenizer(tokenizer_cache);223    }224    vocab_cache     = vocab;225    tokenizer_cache = tokenizer;226 227    return llg_clone_tokenizer(tokenizer_cache);228}229 230llama_sampler * llama_sampler_init_llg(const llama_vocab * vocab, const char * grammar_kind,231                                       const char * grammar_data) {232    auto * ctx = new llama_sampler_llg;233 234    if (grammar_kind != nullptr && grammar_kind[0] != '\0') {235        auto tokenizer = llama_sampler_llg_new_tokenizer(vocab);236        *ctx           = {237            /* .vocab        = */ vocab,238            /* .grammar_kind = */ grammar_kind,239            /* .grammar_data = */ grammar_data,240            /* .tokenizer    = */ tokenizer,241            /* .grammar      = */ llama_sampler_llg_new(tokenizer, grammar_kind, grammar_data),242            /* .llg_res      = */ {},243            /* .has_llg_res  = */ false,244        };245    } else {246        *ctx = {247            /* .vocab        = */ vocab,248            /* .grammar_kind = */ {},249            /* .grammar_data = */ {},250            /* .tokenizer    = */ nullptr,251            /* .grammar      = */ nullptr,252            /* .llg_res      = */ {},253            /* .has_llg_res  = */ false,254        };255    }256 257    return llama_sampler_init(258        /* .iface = */ &llama_sampler_llg_i,259        /* .ctx   = */ ctx260    );261}262 263#else264 265llama_sampler * llama_sampler_init_llg(const llama_vocab *, const char *, const char *) {266    LOG_WRN("llguidance (cmake -DLLAMA_LLGUIDANCE=ON) is not enabled");267    return nullptr;268}269 270#endif  // LLAMA_USE_LLGUIDANCE271