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.
03.1k
1#include "llama-grammar.h"2 3#include "llama-impl.h"4#include "llama-vocab.h"5#include "llama-sampler.h"6 7#include <cmath>8#include <algorithm>9#include <cstdint>10#include <set>11#include <stdexcept>12 13#define MAX_REPETITION_THRESHOLD 200014//15// helpers16//17 18// NOTE: assumes valid utf8 (but checks for overrun)19static std::pair<uint32_t, const char *> decode_utf8(const char * src) {20 static const int lookup[] = { 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 3, 4 };21 uint8_t first_byte = static_cast<uint8_t>(*src);22 uint8_t highbits = first_byte >> 4;23 int len = lookup[highbits];24 uint8_t mask = (1 << (8 - len)) - 1;25 uint32_t value = first_byte & mask;26 const char * end = src + len; // may overrun!27 const char * pos = src + 1;28 for ( ; pos < end && *pos; pos++) {29 value = (value << 6) + (static_cast<uint8_t>(*pos) & 0x3F);30 }31 return std::make_pair(value, pos);32}33 34static std::pair<std::vector<uint32_t>, llama_partial_utf8> decode_utf8(35 const std::string & src,36 llama_partial_utf8 partial_start) {37 static const int lookup[] = { 1, 1, 1, 1, 1, 1, 1, 1, 0, 0, 0, 0, 2, 2, 3, 4 };38 const char * pos = src.c_str();39 std::vector<uint32_t> code_points;40 41 // common english strings have the same number of codepoints and bytes. `+ 1` for the terminating 0.42 code_points.reserve(src.size() + 1);43 uint32_t value = partial_start.value;44 int n_remain = partial_start.n_remain;45 46 // continue previous decode, if applicable47 while (*pos != 0 && n_remain > 0) {48 uint8_t next_byte = static_cast<uint8_t>(*pos);49 if ((next_byte >> 6) != 2) {50 // invalid sequence, abort51 code_points.push_back(0);52 return std::make_pair(std::move(code_points), llama_partial_utf8{ 0, -1 });53 }54 value = (value << 6) + (next_byte & 0x3F);55 ++pos;56 --n_remain;57 }58 59 if (partial_start.n_remain > 0 && n_remain == 0) {60 code_points.push_back(value);61 }62 63 // decode any subsequent utf-8 sequences, which may end in an incomplete one64 while (*pos != 0) {65 uint8_t first_byte = static_cast<uint8_t>(*pos);66 uint8_t highbits = first_byte >> 4;67 n_remain = lookup[highbits] - 1;68 69 if (n_remain < 0) {70 // invalid sequence, abort71 code_points.clear();72 code_points.push_back(0);73 return std::make_pair(std::move(code_points), llama_partial_utf8{ 0, n_remain });74 }75 76 uint8_t mask = (1 << (7 - n_remain)) - 1;77 value = first_byte & mask;78 79 ++pos;80 while (*pos != 0 && n_remain > 0) {81 value = (value << 6) + (static_cast<uint8_t>(*pos) & 0x3F);82 ++pos;83 --n_remain;84 }85 if (n_remain == 0) {86 code_points.push_back(value);87 }88 }89 code_points.push_back(0);90 91 return std::make_pair(std::move(code_points), llama_partial_utf8{ value, n_remain });92}93 94static bool is_digit_char(char c) {95 return '0' <= c && c <= '9';96}97 98static bool is_word_char(char c) {99 return ('a' <= c && c <= 'z') || ('A' <= c && c <= 'Z') || c == '-' || is_digit_char(c);100}101 102static std::pair<uint32_t, const char *> parse_hex(const char * src, int size) {103 const char * pos = src;104 const char * end = src + size;105 uint32_t value = 0;106 for ( ; pos < end && *pos; pos++) {107 value <<= 4;108 char c = *pos;109 if ('a' <= c && c <= 'f') {110 value += c - 'a' + 10;111 } else if ('A' <= c && c <= 'F') {112 value += c - 'A' + 10;113 } else if ('0' <= c && c <= '9') {114 value += c - '0';115 } else {116 break;117 }118 }119 if (pos != end) {120 throw std::runtime_error("expecting " + std::to_string(size) + " hex chars at " + src);121 }122 return std::make_pair(value, pos);123}124 125static const char * parse_space(const char * src, bool newline_ok) {126 const char * pos = src;127 while (*pos == ' ' || *pos == '\t' || *pos == '#' ||128 (newline_ok && (*pos == '\r' || *pos == '\n'))) {129 if (*pos == '#') {130 while (*pos && *pos != '\r' && *pos != '\n') {131 pos++;132 }133 } else {134 pos++;135 }136 }137 return pos;138}139 140static const char * parse_name(const char * src) {141 const char * pos = src;142 while (is_word_char(*pos)) {143 pos++;144 }145 if (pos == src) {146 throw std::runtime_error(std::string("expecting name at ") + src);147 }148 return pos;149}150 151static const char * parse_int(const char * src) {152 const char * pos = src;153 while (is_digit_char(*pos)) {154 pos++;155 }156 if (pos == src) {157 throw std::runtime_error(std::string("expecting integer at ") + src);158 }159 return pos;160}161 162static std::pair<uint32_t, const char *> parse_char(const char * src) {163 if (*src == '\\') {164 switch (src[1]) {165 case 'x': return parse_hex(src + 2, 2);166 case 'u': return parse_hex(src + 2, 4);167 case 'U': return parse_hex(src + 2, 8);168 case 't': return std::make_pair('\t', src + 2);169 case 'r': return std::make_pair('\r', src + 2);170 case 'n': return std::make_pair('\n', src + 2);171 case '\\':172 case '"':173 case '[':174 case ']':175 return std::make_pair(src[1], src + 2);176 default:177 throw std::runtime_error(std::string("unknown escape at ") + src);178 }179 } else if (*src) {180 return decode_utf8(src);181 }182 throw std::runtime_error("unexpected end of input");183}184 185static std::pair<uint32_t, const char *> parse_token(const llama_vocab * vocab, const char * src) {186 const char * pos = src;187 if (*pos != '<') {188 throw std::runtime_error(std::string("expecting '<' at ") + pos);189 }190 pos++;191 192 // Parse <[id]>193 if (*pos == '[') {194 pos++;195 const char * int_end = parse_int(pos);196 uint32_t token_id = std::stoul(std::string(pos, int_end - pos));197 pos = int_end;198 if (*pos != ']') {199 throw std::runtime_error(std::string("expecting ']' at ") + pos);200 }201 pos++;202 if (*pos != '>') {203 throw std::runtime_error(std::string("expecting '>' at ") + pos);204 }205 pos++;206 return std::make_pair(token_id, pos);207 }208 209 if (vocab == nullptr) {210 throw std::runtime_error(std::string("no vocab to parse token at ") + src);211 }212 213 // Parse <token> and tokenize to obtain the token id214 while (*pos != 0 && *pos != '>') {215 pos++;216 }217 if (*pos != '>') {218 throw std::runtime_error(std::string("expecting '>' at ") + pos);219 }220 pos++;221 222 llama_token tokens[2];223 int32_t n_tokens = vocab->tokenize(src, static_cast<int32_t>(pos - src), tokens, 2, false, true);224 if (n_tokens != 1) {225 // must tokenize to exactly 1 token226 throw std::runtime_error("invalid token '" + std::string(src, pos - src) + "'");227 }228 return std::make_pair(tokens[0], pos);229}230 231static void print_grammar_char(FILE * file, uint32_t c) {232 if (0x20 <= c && c <= 0x7f) {233 fprintf(file, "%c", static_cast<char>(c));234 } else {235 // cop out of encoding UTF-8236 fprintf(file, "<U+%04X>", c);237 }238}239 240static bool is_char_element(llama_grammar_element elem) {241 switch (elem.type) {242 case LLAMA_GRETYPE_CHAR: return true;243 case LLAMA_GRETYPE_CHAR_NOT: return true;244 case LLAMA_GRETYPE_CHAR_ALT: return true;245 case LLAMA_GRETYPE_CHAR_RNG_UPPER: return true;246 case LLAMA_GRETYPE_CHAR_ANY: return true;247 default: return false;248 }249}250 251static void print_rule_binary(FILE * file, const llama_grammar_rule & rule) {252 for (auto elem : rule) {253 switch (elem.type) {254 case LLAMA_GRETYPE_END: fprintf(file, "END"); break;255 case LLAMA_GRETYPE_ALT: fprintf(file, "ALT"); break;256 case LLAMA_GRETYPE_RULE_REF: fprintf(file, "RULE_REF"); break;257 case LLAMA_GRETYPE_CHAR: fprintf(file, "CHAR"); break;258 case LLAMA_GRETYPE_CHAR_NOT: fprintf(file, "CHAR_NOT"); break;259 case LLAMA_GRETYPE_CHAR_RNG_UPPER: fprintf(file, "CHAR_RNG_UPPER"); break;260 case LLAMA_GRETYPE_CHAR_ALT: fprintf(file, "CHAR_ALT"); break;261 case LLAMA_GRETYPE_CHAR_ANY: fprintf(file, "CHAR_ANY"); break;262 case LLAMA_GRETYPE_TOKEN: fprintf(file, "TOKEN"); break;263 case LLAMA_GRETYPE_TOKEN_NOT: fprintf(file, "TOKEN_NOT"); break;264 }265 switch (elem.type) {266 case LLAMA_GRETYPE_END:267 case LLAMA_GRETYPE_ALT:268 case LLAMA_GRETYPE_RULE_REF:269 fprintf(file, "(%u) ", elem.value);270 break;271 case LLAMA_GRETYPE_CHAR:272 case LLAMA_GRETYPE_CHAR_NOT:273 case LLAMA_GRETYPE_CHAR_RNG_UPPER:274 case LLAMA_GRETYPE_CHAR_ALT:275 case LLAMA_GRETYPE_CHAR_ANY:276 fprintf(file, "(\"");277 print_grammar_char(file, elem.value);278 fprintf(file, "\") ");279 break;280 case LLAMA_GRETYPE_TOKEN:281 fprintf(file, "<[");282 fprintf(file, "%u", elem.value);283 fprintf(file, "]> ");284 break;285 case LLAMA_GRETYPE_TOKEN_NOT:286 fprintf(file, "!");287 fprintf(file, "<[");288 fprintf(file, "%u", elem.value);289 fprintf(file, "]> ");290 break;291 }292 }293 fprintf(file, "\n");294}295 296static void print_rule(297 FILE * file,298 uint32_t rule_id,299 const llama_grammar_rule & rule,300 const std::map<uint32_t, std::string> & symbol_id_names) {301 if (rule.empty() || rule.back().type != LLAMA_GRETYPE_END) {302 throw std::runtime_error(303 "malformed rule, does not end with LLAMA_GRETYPE_END: " + std::to_string(rule_id));304 }305 fprintf(file, "%s ::= ", symbol_id_names.at(rule_id).c_str());306 for (size_t i = 0, end = rule.size() - 1; i < end; i++) {307 llama_grammar_element elem = rule[i];308 switch (elem.type) {309 case LLAMA_GRETYPE_END:310 throw std::runtime_error(311 "unexpected end of rule: " + std::to_string(rule_id) + "," +312 std::to_string(i));313 case LLAMA_GRETYPE_ALT:314 fprintf(file, "| ");315 break;316 case LLAMA_GRETYPE_RULE_REF:317 fprintf(file, "%s ", symbol_id_names.at(elem.value).c_str());318 break;319 case LLAMA_GRETYPE_CHAR:320 fprintf(file, "[");321 print_grammar_char(file, elem.value);322 break;323 case LLAMA_GRETYPE_CHAR_NOT:324 fprintf(file, "[^");325 print_grammar_char(file, elem.value);326 break;327 case LLAMA_GRETYPE_CHAR_RNG_UPPER:328 if (i == 0 || !is_char_element(rule[i - 1])) {329 throw std::runtime_error(330 "LLAMA_GRETYPE_CHAR_RNG_UPPER without preceding char: " +331 std::to_string(rule_id) + "," + std::to_string(i));332 }333 fprintf(file, "-");334 print_grammar_char(file, elem.value);335 break;336 case LLAMA_GRETYPE_CHAR_ALT:337 if (i == 0 || !is_char_element(rule[i - 1])) {338 throw std::runtime_error(339 "LLAMA_GRETYPE_CHAR_ALT without preceding char: " +340 std::to_string(rule_id) + "," + std::to_string(i));341 }342 print_grammar_char(file, elem.value);343 break;344 case LLAMA_GRETYPE_CHAR_ANY:345 fprintf(file, ".");346 break;347 case LLAMA_GRETYPE_TOKEN:348 fprintf(file, "<[");349 fprintf(file, "%u", elem.value);350 fprintf(file, "]> ");351 break;352 case LLAMA_GRETYPE_TOKEN_NOT:353 fprintf(file, "!");354 fprintf(file, "<[");355 fprintf(file, "%u", elem.value);356 fprintf(file, "]> ");357 break;358 }359 if (is_char_element(elem)) {360 switch (rule[i + 1].type) {361 case LLAMA_GRETYPE_CHAR_ALT:362 case LLAMA_GRETYPE_CHAR_RNG_UPPER:363 case LLAMA_GRETYPE_CHAR_ANY:364 break;365 default:366 fprintf(file, "] ");367 }368 }369 }370 fprintf(file, "\n");371}372 373//374// Regex utilities375//376 377size_t llama_grammar_trigger_pattern::find(const std::string & input) const {378 auto find_start_pos = [](const std::smatch & match) {379 // get from the first matched capturing group to the end of the string380 size_t start = std::string::npos;381 for (auto i = 1u; i < match.size(); i++) {382 if (match.length(i) > 0) {383 start = match.position(i);384 break;385 }386 }387 if (start == std::string::npos) {388 start = match.position(0);389 }390 return start;391 };392 393 if (!pattern.empty() && pattern.front() == '^' && pattern.back() == '$') {394 // match against the entire input395 std::smatch match;396 if (std::regex_match(input, match, regex)) {397 return find_start_pos(match);398 }399 }400 401 // search anywhere402 std::smatch match;403 if (std::regex_search(input, match, regex)) {404 return find_start_pos(match);405 }406 407 return std::string::npos;408}409 410 411//412// implementation413//414 415uint32_t llama_grammar_parser::get_symbol_id(const char * src, size_t len) {416 uint32_t next_id = static_cast<uint32_t>(symbol_ids.size());417 auto result = symbol_ids.emplace(std::string(src, len), next_id);418 return result.first->second;419}420 421uint32_t llama_grammar_parser::generate_symbol_id(const std::string & base_name) {422 uint32_t next_id = static_cast<uint32_t>(symbol_ids.size());423 symbol_ids[base_name + '_' + std::to_string(next_id)] = next_id;424 return next_id;425}426 427void llama_grammar_parser::add_rule(uint32_t rule_id, const llama_grammar_rule & rule) {428 if (rules.size() <= rule_id) {429 rules.resize(rule_id + 1);430 }431 rules[rule_id] = rule;432}433 434const char * llama_grammar_parser::parse_alternates(435 const char * src,436 const std::string & rule_name,437 uint32_t rule_id,438 bool is_nested) {439 llama_grammar_rule rule;440 const char * pos = parse_sequence(src, rule_name, rule, is_nested);441 while (*pos == '|') {442 rule.push_back({LLAMA_GRETYPE_ALT, 0});443 pos = parse_space(pos + 1, true);444 pos = parse_sequence(pos, rule_name, rule, is_nested);445 }446 rule.push_back({LLAMA_GRETYPE_END, 0});447 add_rule(rule_id, rule);448 return pos;449}450 451const char * llama_grammar_parser::parse_sequence(452 const char * src,453 const std::string & rule_name,454 llama_grammar_rule & rule,455 bool is_nested) {456 size_t last_sym_start = rule.size();457 const char * pos = src;458 uint64_t n_prev_rules = 1;459 460 // use UINT64_MAX as the empty value because we aligned to the proper uint64_t type so -1 can't be used461 // (though it's technically the same as -1 now)462 auto handle_repetitions = [&](uint64_t min_times, uint64_t max_times) {463 bool no_max = max_times == UINT64_MAX;464 if (last_sym_start == rule.size()) {465 throw std::runtime_error(std::string("expecting preceding item to */+/?/{ at ") + pos);466 }467 468 // apply transformation to previous symbol (last_sym_start to end) according to469 // the following rewrite rules:470 // S{m,n} --> S S S (m times) S'(n-m)471 // S'(x) ::= S S'(x-1) |472 // (... n-m definitions of these S' rules ...)473 // S'(1) ::= S |474 // S{m,} --> S S S (m times) S'475 // S' ::= S S' |476 // S* --> S{0,}477 // --> S' ::= S S' |478 // S+ --> S{1,}479 // --> S S'480 // S' ::= S S' |481 // S? --> S{0,1}482 // --> S'483 // S' ::= S |484 485 llama_grammar_rule prev_rule(rule.begin() + last_sym_start, rule.end());486 // Calculate the total number of rules that will be generated by this repetition487 uint64_t total_rules = 1; // Start with 1 for the original rule488 if (!no_max && max_times > 0) {489 total_rules = max_times;490 } else if (min_times > 0) {491 total_rules = min_times;492 }493 494 if (n_prev_rules * total_rules >= MAX_REPETITION_THRESHOLD) {495 throw std::runtime_error("number of rules that are going to be repeated multiplied by the new repetition exceeds sane defaults, please reduce the number of repetitions or rule complexity");496 }497 498 if (min_times == 0) {499 rule.resize(last_sym_start);500 } else {501 // Repeat the previous elements (min_times - 1) times502 for (uint64_t i = 1; i < min_times; i++) {503 rule.insert(rule.end(), prev_rule.begin(), prev_rule.end());504 }505 }506 507 uint32_t last_rec_rule_id = 0;508 auto n_opt = no_max ? 1 : max_times - min_times;509 510 llama_grammar_rule rec_rule(prev_rule);511 for (uint64_t i = 0; i < n_opt; i++) {512 rec_rule.resize(prev_rule.size());513 uint32_t rec_rule_id = generate_symbol_id( rule_name);514 if (i > 0 || no_max) {515 rec_rule.push_back({LLAMA_GRETYPE_RULE_REF, no_max ? rec_rule_id : last_rec_rule_id});516 }517 rec_rule.push_back({LLAMA_GRETYPE_ALT, 0});518 rec_rule.push_back({LLAMA_GRETYPE_END, 0});519 add_rule( rec_rule_id, rec_rule);520 last_rec_rule_id = rec_rule_id;521 }522 if (n_opt > 0) {523 rule.push_back({LLAMA_GRETYPE_RULE_REF, last_rec_rule_id});524 }525 n_prev_rules *= total_rules;526 GGML_ASSERT(n_prev_rules >= 1);527 };528 529 while (*pos) {530 if (*pos == '"') { // literal string531 pos++;532 last_sym_start = rule.size();533 n_prev_rules = 1;534 while (*pos != '"') {535 if (!*pos) {536 throw std::runtime_error("unexpected end of input");537 }538 auto char_pair = parse_char(pos);539 pos = char_pair.second;540 rule.push_back({LLAMA_GRETYPE_CHAR, char_pair.first});541 }542 pos = parse_space(pos + 1, is_nested);543 } else if (*pos == '[') { // char range(s)544 pos++;545 enum llama_gretype start_type = LLAMA_GRETYPE_CHAR;546 if (*pos == '^') {547 pos++;548 start_type = LLAMA_GRETYPE_CHAR_NOT;549 }550 last_sym_start = rule.size();551 n_prev_rules = 1;552 while (*pos != ']') {553 if (!*pos) {554 throw std::runtime_error("unexpected end of input");555 }556 auto char_pair = parse_char(pos);557 pos = char_pair.second;558 enum llama_gretype type = last_sym_start < rule.size()559 ? LLAMA_GRETYPE_CHAR_ALT560 : start_type;561 562 rule.push_back({type, char_pair.first});563 if (pos[0] == '-' && pos[1] != ']') {564 if (!pos[1]) {565 throw std::runtime_error("unexpected end of input");566 }567 auto endchar_pair = parse_char(pos + 1);568 pos = endchar_pair.second;569 rule.push_back({LLAMA_GRETYPE_CHAR_RNG_UPPER, endchar_pair.first});570 }571 }572 pos = parse_space(pos + 1, is_nested);573 } else if (*pos == '<' || *pos == '!') { // token574 auto type = LLAMA_GRETYPE_TOKEN;575 if (*pos == '!') { // token inverse576 type = LLAMA_GRETYPE_TOKEN_NOT;577 pos++;578 }579 auto token_pair = parse_token(vocab, pos);580 const char * token_end = token_pair.second;581 last_sym_start = rule.size();582 n_prev_rules = 1;583 rule.push_back({type, token_pair.first});584 pos = parse_space(token_end, is_nested);585 } else if (is_word_char(*pos)) { // rule reference586 const char * name_end = parse_name(pos);587 uint32_t ref_rule_id = get_symbol_id(pos, name_end - pos);588 pos = parse_space(name_end, is_nested);589 last_sym_start = rule.size();590 n_prev_rules = 1;591 rule.push_back({LLAMA_GRETYPE_RULE_REF, ref_rule_id});592 } else if (*pos == '(') { // grouping593 // parse nested alternates into synthesized rule594 pos = parse_space(pos + 1, true);595 uint32_t n_rules_before = symbol_ids.size();596 uint32_t sub_rule_id = generate_symbol_id(rule_name);597 pos = parse_alternates(pos, rule_name, sub_rule_id, true);598 n_prev_rules = std::max(1u, (uint32_t)symbol_ids.size() - n_rules_before);599 last_sym_start = rule.size();600 // output reference to synthesized rule601 rule.push_back({LLAMA_GRETYPE_RULE_REF, sub_rule_id});602 if (*pos != ')') {603 throw std::runtime_error(std::string("expecting ')' at ") + pos);604 }605 pos = parse_space(pos + 1, is_nested);606 } else if (*pos == '.') { // any char607 last_sym_start = rule.size();608 n_prev_rules = 1;609 rule.push_back({LLAMA_GRETYPE_CHAR_ANY, 0});610 pos = parse_space(pos + 1, is_nested);611 } else if (*pos == '*') {612 pos = parse_space(pos + 1, is_nested);613 handle_repetitions(0, -1);614 } else if (*pos == '+') {615 pos = parse_space(pos + 1, is_nested);616 handle_repetitions(1, -1);617 } else if (*pos == '?') {618 pos = parse_space(pos + 1, is_nested);619 handle_repetitions(0, 1);620 } else if (*pos == '{') {621 pos = parse_space(pos + 1, is_nested);622 623 if (!is_digit_char(*pos)) {624 throw std::runtime_error(std::string("expecting an int at ") + pos);625 }626 const char * int_end = parse_int(pos);627 uint64_t min_times = std::stoull(std::string(pos, int_end - pos));628 pos = parse_space(int_end, is_nested);629 630 uint64_t max_times = UINT64_MAX; // default: no max limit631 632 if (*pos == '}') {633 max_times = min_times;634 pos = parse_space(pos + 1, is_nested);635 } else if (*pos == ',') {636 pos = parse_space(pos + 1, is_nested);637 638 if (is_digit_char(*pos)) {639 const char * int_end = parse_int(pos);640 max_times = std::stoull(std::string(pos, int_end - pos));641 pos = parse_space(int_end, is_nested);642 }643 644 if (*pos != '}') {645 throw std::runtime_error(std::string("expecting '}' at ") + pos);646 }647 pos = parse_space(pos + 1, is_nested);648 } else {649 throw std::runtime_error(std::string("expecting ',' at ") + pos);650 }651 if (min_times > MAX_REPETITION_THRESHOLD) {652 throw std::runtime_error(std::string("number of repetitions exceeds sane defaults, please reduce the number of repetitions"));653 }654 if (max_times != UINT64_MAX && max_times > MAX_REPETITION_THRESHOLD) {655 max_times = UINT64_MAX;656 }657 handle_repetitions(min_times, max_times);658 } else {659 break;660 }661 }662 return pos;663}664 665const char * llama_grammar_parser::parse_rule(const char * src) {666 const char * name_end = parse_name(src);667 const char * pos = parse_space(name_end, false);668 size_t name_len = name_end - src;669 uint32_t rule_id = get_symbol_id(src, name_len);670 const std::string name(src, name_len);671 672 if (!(pos[0] == ':' && pos[1] == ':' && pos[2] == '=')) {673 throw std::runtime_error(std::string("expecting ::= at ") + pos);674 }675 pos = parse_space(pos + 3, true);676 677 pos = parse_alternates(pos, name, rule_id, false);678 679 if (*pos == '\r') {680 pos += pos[1] == '\n' ? 2 : 1;681 } else if (*pos == '\n') {682 pos++;683 } else if (*pos) {684 throw std::runtime_error(std::string("expecting newline or end at ") + pos);685 }686 return parse_space(pos, true);687}688 689bool llama_grammar_parser::parse(const char * src) {690 try {691 const char * pos = parse_space(src, true);692 while (*pos) {693 pos = parse_rule(pos);694 }695 // Validate the state to ensure that all rules are defined696 for (const auto & rule : rules) {697 if (rule.empty()) {698 throw std::runtime_error("Undefined rule");699 }700 for (const auto & elem : rule) {701 if (elem.type == LLAMA_GRETYPE_RULE_REF) {702 // Ensure that the rule at that location exists703 if (elem.value >= rules.size() || rules[elem.value].empty()) {704 // Get the name of the rule that is missing705 for (const auto & kv : symbol_ids) {706 if (kv.second == elem.value) {707 throw std::runtime_error("Undefined rule identifier '" + kv.first + "'");708 }709 }710 }711 }712 }713 }714 } catch (const std::exception & err) {715 fprintf(stderr, "%s: error parsing grammar: %s\n\n%s\n", __func__, err.what(), src);716 rules.clear();717 return false;718 }719 720 return true;721}722 723void llama_grammar_parser::print(FILE * file) {724 try {725 std::map<uint32_t, std::string> symbol_id_names;726 for (const auto & kv : symbol_ids) {727 symbol_id_names[kv.second] = kv.first;728 }729 for (size_t i = 0, end = rules.size(); i < end; i++) {730 // fprintf(file, "%zu: ", i);731 // print_rule_binary(file, rules[i]);732 print_rule(file, uint32_t(i), rules[i], symbol_id_names);733 // fprintf(file, "\n");734 }735 } catch (const std::exception & err) {736 fprintf(stderr, "\n%s: error printing grammar: %s\n", __func__, err.what());737 }738}739 740llama_grammar_stack llama_grammar_parser::c_rules() const {741 llama_grammar_stack ret;742 ret.reserve(rules.size());743 for (const auto & rule : rules) {744 ret.push_back(rule.data());745 }746 return ret;747}748 749// returns true iff pos points to the end of one of the definitions of a rule750static bool llama_grammar_is_end_of_sequence(const llama_grammar_element * pos) {751 switch (pos->type) {752 case LLAMA_GRETYPE_END: return true; // NOLINT753 case LLAMA_GRETYPE_ALT: return true; // NOLINT754 default: return false;755 }756}757 758// returns true iff chr satisfies the char range at pos (regular or inverse range)759// asserts that pos is pointing to a char range element760static std::pair<bool, const llama_grammar_element *> llama_grammar_match_char(761 const llama_grammar_element * pos,762 const uint32_t chr) {763 bool found = false;764 bool is_positive_char = pos->type == LLAMA_GRETYPE_CHAR || pos->type == LLAMA_GRETYPE_CHAR_ANY;765 766 GGML_ASSERT(is_positive_char || pos->type == LLAMA_GRETYPE_CHAR_NOT); // NOLINT767 768 do {769 if (pos[1].type == LLAMA_GRETYPE_CHAR_RNG_UPPER) {770 // inclusive range, e.g. [a-z]771 found = found || (pos->value <= chr && chr <= pos[1].value);772 pos += 2;773 } else if (pos->type == LLAMA_GRETYPE_CHAR_ANY) {774 // Any character matches "."775 found = true;776 pos += 1;777 } else {778 // exact char match, e.g. [a] or "a"779 found = found || pos->value == chr;780 pos += 1;781 }782 } while (pos->type == LLAMA_GRETYPE_CHAR_ALT);783 784 return std::make_pair(found == is_positive_char, pos);785}786 787// returns true iff some continuation of the given partial UTF-8 sequence could satisfy the char788// range at pos (regular or inverse range)789// asserts that pos is pointing to a char range element790static bool llama_grammar_match_partial_char(791 const llama_grammar_element * pos,792 const llama_partial_utf8 partial_utf8) {793 bool is_positive_char = pos->type == LLAMA_GRETYPE_CHAR || pos->type == LLAMA_GRETYPE_CHAR_ANY;794 GGML_ASSERT(is_positive_char || pos->type == LLAMA_GRETYPE_CHAR_NOT);795 796 uint32_t partial_value = partial_utf8.value;797 int n_remain = partial_utf8.n_remain;798 799 // invalid sequence or 7-bit char split across 2 bytes (overlong)800 if (n_remain < 0 || (n_remain == 1 && partial_value < 2)) {801 return false;802 }803 804 // range of possible code points this partial UTF-8 sequence could complete to805 uint32_t low = partial_value << (n_remain * 6);806 uint32_t high = low | ((1 << (n_remain * 6)) - 1);807 808 if (low == 0) {809 if (n_remain == 2) {810 low = 1 << 11;811 } else if (n_remain == 3) {812 low = 1 << 16;813 }814 }815 816 do {817 if (pos[1].type == LLAMA_GRETYPE_CHAR_RNG_UPPER) {818 // inclusive range, e.g. [a-z]819 if (pos->value <= high && low <= pos[1].value) {820 return is_positive_char;821 }822 pos += 2;823 } else if (pos->type == LLAMA_GRETYPE_CHAR_ANY) {824 // Any character matches "."825 return true;826 } else {827 // exact char match, e.g. [a] or "a"828 if (low <= pos->value && pos->value <= high) {829 return is_positive_char;830 }831 pos += 1;832 }833 } while (pos->type == LLAMA_GRETYPE_CHAR_ALT);834 835 return !is_positive_char;836}837 838// returns true iff token matches the rule at pos (regular or inverse)839// asserts that pos is pointing to a token element840static bool llama_grammar_match_token(841 const llama_grammar_element * pos,842 const llama_token token) {843 GGML_ASSERT(pos->type == LLAMA_GRETYPE_TOKEN || pos->type == LLAMA_GRETYPE_TOKEN_NOT);844 if (pos->type == LLAMA_GRETYPE_TOKEN) {845 return pos->value == static_cast<uint32_t>(token);846 }847 if (pos->type == LLAMA_GRETYPE_TOKEN_NOT) {848 return pos->value != static_cast<uint32_t>(token);849 }850 return false;851}852 853// transforms a grammar pushdown stack into N possible stacks, all ending854// at a character range (terminal element)855static void llama_grammar_advance_stack(856 const llama_grammar_rules & rules,857 const llama_grammar_stack & stack,858 llama_grammar_stacks & new_stacks) {859 std::vector<llama_grammar_stack> todo;860 todo.push_back(stack);861 862 auto stack_cmp = [](const llama_grammar_stack & a, const llama_grammar_stack & b) {863 return std::lexicographical_compare(a.begin(), a.end(), b.begin(), b.end(),864 [](const llama_grammar_element * pa, const llama_grammar_element * pb) {865 return pa < pb; // Compare pointer addresses866 }867 );868 };869 870 std::set<llama_grammar_stack, decltype(stack_cmp)> seen(stack_cmp);871 872 while (!todo.empty()) {873 llama_grammar_stack curr_stack = std::move(todo.back());874 todo.pop_back();875 876 if (seen.find( curr_stack) != seen.end()) {877 continue;878 }879 seen.insert(curr_stack);880 881 if (curr_stack.empty()) {882 if (std::find(new_stacks.begin(), new_stacks.end(), curr_stack) == new_stacks.end()) {883 new_stacks.emplace_back(std::move(curr_stack));884 }885 continue;886 }887 888 const llama_grammar_element * pos = curr_stack.back();889 890 switch (pos->type) {891 case LLAMA_GRETYPE_RULE_REF: {892 const size_t rule_id = static_cast<size_t>(pos->value);893 const llama_grammar_element * subpos = rules[rule_id].data();894 do {895 // init new stack without the top (pos)896 llama_grammar_stack next_stack(curr_stack.begin(), curr_stack.end() - 1);897 if (!llama_grammar_is_end_of_sequence(pos + 1)) {898 // if this rule ref is followed by another element, add that to stack899 next_stack.push_back(pos + 1);900 }901 if (!llama_grammar_is_end_of_sequence(subpos)) {902 // if alternate is nonempty, add to stack903 next_stack.push_back(subpos);904 }905 todo.push_back(std::move(next_stack));906 while (!llama_grammar_is_end_of_sequence(subpos)) {907 // scan to end of alternate def908 subpos++;909 }910 if (subpos->type == LLAMA_GRETYPE_ALT) {911 // there's another alternate def of this rule to process912 subpos++;913 } else {914 break;915 }916 } while (true);917 break;918 }919 case LLAMA_GRETYPE_CHAR:920 case LLAMA_GRETYPE_CHAR_NOT:921 case LLAMA_GRETYPE_CHAR_ANY:922 case LLAMA_GRETYPE_TOKEN:923 case LLAMA_GRETYPE_TOKEN_NOT:924 if (std::find(new_stacks.begin(), new_stacks.end(), curr_stack) == new_stacks.end()) {925 // only add the stack if it's not a duplicate of one we already have926 new_stacks.emplace_back(std::move(curr_stack));927 }928 break;929 default:930 // end of alternate (LLAMA_GRETYPE_END, LLAMA_GRETYPE_ALT) or middle of char range931 // (LLAMA_GRETYPE_CHAR_ALT, LLAMA_GRETYPE_CHAR_RNG_UPPER); stack should never be left on932 // those933 GGML_ABORT("fatal error");934 }935 }936}937 938static llama_grammar_candidates llama_grammar_reject_candidates(939 const llama_grammar_rules & rules,940 const llama_grammar_stacks & stacks,941 const llama_grammar_candidates & candidates) {942 GGML_ASSERT(!stacks.empty()); // REVIEW943 944 if (candidates.empty()) {945 return {};946 }947 948 auto rejects = llama_grammar_reject_candidates_for_stack(rules, stacks.front(), candidates);949 950 for (size_t i = 1, size = stacks.size(); i < size; ++i) {951 rejects = llama_grammar_reject_candidates_for_stack(rules, stacks[i], rejects);952 }953 954 return rejects;955}956 957static bool llama_grammar_detect_left_recursion(958 const llama_grammar_rules & rules,959 size_t rule_index,960 std::vector<bool> * rules_visited,961 std::vector<bool> * rules_in_progress,962 std::vector<bool> * rules_may_be_empty) {963 if ((*rules_in_progress)[rule_index]) {964 return true;965 }966 967 (*rules_in_progress)[rule_index] = true;968 969 const llama_grammar_rule & rule = rules[rule_index];970 971 // First check if the rule might produce the empty string. This could be done combined with the second972 // step but it's more readable as two steps.973 bool at_rule_start = true;974 for (size_t i = 0; i < rule.size(); i++) {975 if (llama_grammar_is_end_of_sequence(&rule[i])) {976 if (at_rule_start) {977 (*rules_may_be_empty)[rule_index] = true;978 break;979 }980 at_rule_start = true;981 } else {982 at_rule_start = false;983 }984 }985 986 // Second, recurse into leftmost nonterminals (or next-leftmost as long as the previous nonterminal may987 // be empty)988 bool recurse_into_nonterminal = true;989 for (size_t i = 0; i < rule.size(); i++) {990 if (rule[i].type == LLAMA_GRETYPE_RULE_REF && recurse_into_nonterminal) {991 if (llama_grammar_detect_left_recursion(rules, (size_t)rule[i].value, rules_visited, rules_in_progress, rules_may_be_empty)) {992 return true;993 }994 if (!((*rules_may_be_empty)[(size_t)rule[i].value])) {995 recurse_into_nonterminal = false;996 }997 } else if (llama_grammar_is_end_of_sequence(&rule[i])) {998 recurse_into_nonterminal = true;999 } else {1000 recurse_into_nonterminal = false;1001 }1002 }1003 1004 (*rules_in_progress)[rule_index] = false;1005 (*rules_visited)[rule_index] = true;1006 1007 return false;1008}1009 1010const llama_grammar_rules & llama_grammar_get_rules(const struct llama_grammar * grammar) {1011 return grammar->rules;1012}1013 1014llama_grammar_stacks & llama_grammar_get_stacks(struct llama_grammar * grammar) {1015 return grammar->stacks;1016}1017 1018static void llama_grammar_accept_chr(1019 struct llama_grammar & grammar,1020 const llama_grammar_stack & stack,1021 uint32_t chr,1022 llama_grammar_stacks & new_stacks) {1023 if (stack.empty()) {1024 return;1025 }1026 1027 const llama_grammar_element * pos = stack.back();1028 1029 // ignore if this turns into a token1030 if (pos->type == LLAMA_GRETYPE_TOKEN || pos->type == LLAMA_GRETYPE_TOKEN_NOT) {1031 return;1032 }1033 1034 auto match = llama_grammar_match_char(pos, chr);1035 if (match.first) {1036 llama_grammar_stack new_stack(stack.begin(), stack.end() - 1);1037 if (!llama_grammar_is_end_of_sequence(match.second)) {1038 new_stack.push_back(match.second);1039 }1040 llama_grammar_advance_stack(grammar.rules, new_stack, new_stacks);1041 }1042}1043 1044void llama_grammar_accept(struct llama_grammar * grammar, uint32_t chr) {1045 llama_grammar_stacks stacks_new;1046 stacks_new.reserve(grammar->stacks.size());1047 1048 for (const auto & stack : grammar->stacks) {1049 llama_grammar_accept_chr(*grammar, stack, chr, stacks_new);1050 }1051 1052 grammar->stacks = std::move(stacks_new);1053}1054 1055llama_grammar_candidates llama_grammar_reject_candidates_for_stack(1056 const llama_grammar_rules & rules,1057 const llama_grammar_stack & stack,1058 const llama_grammar_candidates & candidates) {1059 1060 llama_grammar_candidates rejects;1061 rejects.reserve(candidates.size());1062 1063 if (stack.empty()) {1064 for (const auto & tok : candidates) {1065 if (*tok.code_points != 0 || tok.partial_utf8.n_remain != 0) {1066 rejects.push_back(tok);1067 }1068 }1069 return rejects;1070 }1071 1072 const llama_grammar_element * stack_pos = stack.back();1073 1074 // if the top of the stack is a token rule, then we only need to check the token id1075 if (stack_pos->type == LLAMA_GRETYPE_TOKEN || stack_pos->type == LLAMA_GRETYPE_TOKEN_NOT) {1076 for (const auto & tok : candidates) {1077 if (*tok.code_points == 0) {1078 // reached the end of a token consumed by char rules, reject iff it ended1079 // in a partial response1080 if (tok.partial_utf8.n_remain != 0) {1081 rejects.push_back(tok);1082 }1083 } else if (!llama_grammar_match_token(stack_pos, tok.id)) {1084 rejects.push_back(tok);1085 }1086 }1087 return rejects;1088 }1089 1090 llama_grammar_candidates next_candidates;1091 next_candidates.reserve(candidates.size());1092 1093 for (const auto & tok : candidates) {1094 if (*tok.code_points == 0) {1095 // reached end of full codepoints in token, reject iff it ended in a partial sequence1096 // that cannot satisfy this position in grammar1097 if (tok.partial_utf8.n_remain != 0 &&1098 !llama_grammar_match_partial_char(stack_pos, tok.partial_utf8)) {1099 rejects.push_back(tok);1100 }1101 } else if (llama_grammar_match_char(stack_pos, *tok.code_points).first) {1102 next_candidates.push_back({ tok.index, tok.code_points + 1, tok.partial_utf8, tok.id });1103 } else {1104 rejects.push_back(tok);1105 }1106 }1107 1108 const auto * stack_pos_after = llama_grammar_match_char(stack_pos, 0).second;1109 1110 // update top of stack to next element, if any1111 llama_grammar_stack stack_after(stack.begin(), stack.end() - 1);1112 if (!llama_grammar_is_end_of_sequence(stack_pos_after)) {1113 stack_after.push_back(stack_pos_after);1114 }1115 llama_grammar_stacks next_stacks;1116 llama_grammar_advance_stack(rules, stack_after, next_stacks);1117 1118 auto next_rejects = llama_grammar_reject_candidates(rules, next_stacks, next_candidates);1119 for (const auto & tok : next_rejects) {1120 rejects.push_back({ tok.index, tok.code_points - 1, tok.partial_utf8, tok.id });1121 }1122 1123 return rejects;1124}1125 1126////////////////////1127 1128struct llama_grammar * llama_grammar_init_impl(1129 const struct llama_vocab * vocab,1130 const llama_grammar_element ** rules,1131 size_t n_rules,1132 size_t start_rule_index) {1133 const llama_grammar_element * pos;1134 1135 // copy rule definitions into vectors1136 llama_grammar_rules vec_rules(n_rules);1137 for (size_t i = 0; i < n_rules; i++) {1138 for (pos = rules[i]; pos->type != LLAMA_GRETYPE_END; pos++) {1139 vec_rules[i].push_back(*pos);1140 }1141 vec_rules[i].push_back({LLAMA_GRETYPE_END, 0});1142 }1143 1144 // Validate that all rule references point to valid rules1145 for (size_t i = 0; i < n_rules; i++) {1146 for (const auto & elem : vec_rules[i]) {1147 if (elem.type == LLAMA_GRETYPE_RULE_REF) {1148 if (elem.value >= n_rules || vec_rules[elem.value].empty()) {1149 LLAMA_LOG_ERROR("invalid grammar: rule %zu references undefined rule %u\n", i, elem.value);1150 return nullptr;1151 }1152 }1153 }1154 }1155 1156 // Check for left recursion1157 std::vector<bool> rules_visited(n_rules);1158 std::vector<bool> rules_in_progress(n_rules);1159 std::vector<bool> rules_may_be_empty(n_rules);1160 for (size_t i = 0; i < n_rules; i++) {1161 if (rules_visited[i]) {1162 continue;1163 }1164 if (llama_grammar_detect_left_recursion(vec_rules, i, &rules_visited, &rules_in_progress, &rules_may_be_empty)) {1165 LLAMA_LOG_ERROR("unsupported grammar, left recursion detected for nonterminal at index %zu", i);1166 return nullptr;1167 }1168 }1169 1170 // loop over alternates of start rule to build initial stacks1171 llama_grammar_stacks stacks;1172 pos = vec_rules[start_rule_index].data();1173 do {1174 llama_grammar_stack stack;1175 if (!llama_grammar_is_end_of_sequence(pos)) {1176 // if alternate is nonempty, add to stack1177 stack.push_back(pos);1178 }1179 llama_grammar_advance_stack(vec_rules, stack, stacks);1180 while (!llama_grammar_is_end_of_sequence(pos)) {1181 // scan to end of alternate def1182 pos++;1183 }1184 if (pos->type == LLAMA_GRETYPE_ALT) {1185 // there's another alternate def of this rule to process1186 pos++;1187 } else {1188 break;1189 }1190 } while (true);1191 1192 // Important: vec_rules has to be moved here, not copied, because stacks contains1193 // pointers to elements of vec_rules. If vec_rules were copied into llama_grammar1194 // then the pointers would be invalidated when the local vec_rules goes out of scope.1195 return new llama_grammar {1196 vocab,1197 std::move(vec_rules),1198 std::move(stacks),1199 /* .partial_utf8 = */ {},1200 /* .lazy = */ false,