KBaba7/llama.cpp
0
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 