echodict/llama.cpp
version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786
0479
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 bool has_max = max_times != UINT64_MAX;652 if (min_times > MAX_REPETITION_THRESHOLD || (has_max && max_times > MAX_REPETITION_THRESHOLD)) {653 throw std::runtime_error(std::string("number of repetitions exceeds sane defaults, please reduce the number of repetitions"));654 }655 handle_repetitions(min_times, max_times);656 } else {657 break;658 }659 }660 return pos;661}662 663const char * llama_grammar_parser::parse_rule(const char * src) {664 const char * name_end = parse_name(src);665 const char * pos = parse_space(name_end, false);666 size_t name_len = name_end - src;667 uint32_t rule_id = get_symbol_id(src, name_len);668 const std::string name(src, name_len);669 670 if (!(pos[0] == ':' && pos[1] == ':' && pos[2] == '=')) {671 throw std::runtime_error(std::string("expecting ::= at ") + pos);672 }673 pos = parse_space(pos + 3, true);674 675 pos = parse_alternates(pos, name, rule_id, false);676 677 if (*pos == '\r') {678 pos += pos[1] == '\n' ? 2 : 1;679 } else if (*pos == '\n') {680 pos++;681 } else if (*pos) {682 throw std::runtime_error(std::string("expecting newline or end at ") + pos);683 }684 return parse_space(pos, true);685}686 687bool llama_grammar_parser::parse(const char * src) {688 try {689 const char * pos = parse_space(src, true);690 while (*pos) {691 pos = parse_rule(pos);692 }693 // Validate the state to ensure that all rules are defined694 for (const auto & rule : rules) {695 if (rule.empty()) {696 throw std::runtime_error("Undefined rule");697 }698 for (const auto & elem : rule) {699 if (elem.type == LLAMA_GRETYPE_RULE_REF) {700 // Ensure that the rule at that location exists701 if (elem.value >= rules.size() || rules[elem.value].empty()) {702 // Get the name of the rule that is missing703 for (const auto & kv : symbol_ids) {704 if (kv.second == elem.value) {705 throw std::runtime_error("Undefined rule identifier '" + kv.first + "'");706 }707 }708 }709 }710 }711 }712 } catch (const std::exception & err) {713 fprintf(stderr, "%s: error parsing grammar: %s\n\n%s\n", __func__, err.what(), src);714 rules.clear();715 return false;716 }717 718 return true;719}720 721void llama_grammar_parser::print(FILE * file) {722 try {723 std::map<uint32_t, std::string> symbol_id_names;724 for (const auto & kv : symbol_ids) {725 symbol_id_names[kv.second] = kv.first;726 }727 for (size_t i = 0, end = rules.size(); i < end; i++) {728 // fprintf(file, "%zu: ", i);729 // print_rule_binary(file, rules[i]);730 print_rule(file, uint32_t(i), rules[i], symbol_id_names);731 // fprintf(file, "\n");732 }733 } catch (const std::exception & err) {734 fprintf(stderr, "\n%s: error printing grammar: %s\n", __func__, err.what());735 }736}737 738llama_grammar_stack llama_grammar_parser::c_rules() const {739 llama_grammar_stack ret;740 ret.reserve(rules.size());741 for (const auto & rule : rules) {742 ret.push_back(rule.data());743 }744 return ret;745}746 747// returns true iff pos points to the end of one of the definitions of a rule748static bool llama_grammar_is_end_of_sequence(const llama_grammar_element * pos) {749 switch (pos->type) {750 case LLAMA_GRETYPE_END: return true; // NOLINT751 case LLAMA_GRETYPE_ALT: return true; // NOLINT752 default: return false;753 }754}755 756// returns true iff chr satisfies the char range at pos (regular or inverse range)757// asserts that pos is pointing to a char range element758static std::pair<bool, const llama_grammar_element *> llama_grammar_match_char(759 const llama_grammar_element * pos,760 const uint32_t chr) {761 bool found = false;762 bool is_positive_char = pos->type == LLAMA_GRETYPE_CHAR || pos->type == LLAMA_GRETYPE_CHAR_ANY;763 764 GGML_ASSERT(is_positive_char || pos->type == LLAMA_GRETYPE_CHAR_NOT); // NOLINT765 766 do {767 if (pos[1].type == LLAMA_GRETYPE_CHAR_RNG_UPPER) {768 // inclusive range, e.g. [a-z]769 found = found || (pos->value <= chr && chr <= pos[1].value);770 pos += 2;771 } else if (pos->type == LLAMA_GRETYPE_CHAR_ANY) {772 // Any character matches "."773 found = true;774 pos += 1;775 } else {776 // exact char match, e.g. [a] or "a"777 found = found || pos->value == chr;778 pos += 1;779 }780 } while (pos->type == LLAMA_GRETYPE_CHAR_ALT);781 782 return std::make_pair(found == is_positive_char, pos);783}784 785// returns true iff some continuation of the given partial UTF-8 sequence could satisfy the char786// range at pos (regular or inverse range)787// asserts that pos is pointing to a char range element788static bool llama_grammar_match_partial_char(789 const llama_grammar_element * pos,790 const llama_partial_utf8 partial_utf8) {791 bool is_positive_char = pos->type == LLAMA_GRETYPE_CHAR || pos->type == LLAMA_GRETYPE_CHAR_ANY;792 GGML_ASSERT(is_positive_char || pos->type == LLAMA_GRETYPE_CHAR_NOT);793 794 uint32_t partial_value = partial_utf8.value;795 int n_remain = partial_utf8.n_remain;796 797 // invalid sequence or 7-bit char split across 2 bytes (overlong)798 if (n_remain < 0 || (n_remain == 1 && partial_value < 2)) {799 return false;800 }801 802 // range of possible code points this partial UTF-8 sequence could complete to803 uint32_t low = partial_value << (n_remain * 6);804 uint32_t high = low | ((1 << (n_remain * 6)) - 1);805 806 if (low == 0) {807 if (n_remain == 2) {808 low = 1 << 11;809 } else if (n_remain == 3) {810 low = 1 << 16;811 }812 }813 814 do {815 if (pos[1].type == LLAMA_GRETYPE_CHAR_RNG_UPPER) {816 // inclusive range, e.g. [a-z]817 if (pos->value <= high && low <= pos[1].value) {818 return is_positive_char;819 }820 pos += 2;821 } else if (pos->type == LLAMA_GRETYPE_CHAR_ANY) {822 // Any character matches "."823 return true;824 } else {825 // exact char match, e.g. [a] or "a"826 if (low <= pos->value && pos->value <= high) {827 return is_positive_char;828 }829 pos += 1;830 }831 } while (pos->type == LLAMA_GRETYPE_CHAR_ALT);832 833 return !is_positive_char;834}835 836// returns true iff token matches the rule at pos (regular or inverse)837// asserts that pos is pointing to a token element838static bool llama_grammar_match_token(839 const llama_grammar_element * pos,840 const llama_token token) {841 GGML_ASSERT(pos->type == LLAMA_GRETYPE_TOKEN || pos->type == LLAMA_GRETYPE_TOKEN_NOT);842 if (pos->type == LLAMA_GRETYPE_TOKEN) {843 return pos->value == static_cast<uint32_t>(token);844 }845 if (pos->type == LLAMA_GRETYPE_TOKEN_NOT) {846 return pos->value != static_cast<uint32_t>(token);847 }848 return false;849}850 851// transforms a grammar pushdown stack into N possible stacks, all ending852// at a character range (terminal element)853static void llama_grammar_advance_stack(854 const llama_grammar_rules & rules,855 const llama_grammar_stack & stack,856 llama_grammar_stacks & new_stacks) {857 std::vector<llama_grammar_stack> todo;858 todo.push_back(stack);859 860 auto stack_cmp = [](const llama_grammar_stack & a, const llama_grammar_stack & b) {861 return std::lexicographical_compare(a.begin(), a.end(), b.begin(), b.end(),862 [](const llama_grammar_element * pa, const llama_grammar_element * pb) {863 return pa < pb; // Compare pointer addresses864 }865 );866 };867 868 std::set<llama_grammar_stack, decltype(stack_cmp)> seen(stack_cmp);869 870 while (!todo.empty()) {871 llama_grammar_stack curr_stack = std::move(todo.back());872 todo.pop_back();873 874 if (seen.find( curr_stack) != seen.end()) {875 continue;876 }877 seen.insert(curr_stack);878 879 if (curr_stack.empty()) {880 if (std::find(new_stacks.begin(), new_stacks.end(), curr_stack) == new_stacks.end()) {881 new_stacks.emplace_back(std::move(curr_stack));882 }883 continue;884 }885 886 const llama_grammar_element * pos = curr_stack.back();887 888 switch (pos->type) {889 case LLAMA_GRETYPE_RULE_REF: {890 const size_t rule_id = static_cast<size_t>(pos->value);891 const llama_grammar_element * subpos = rules[rule_id].data();892 do {893 // init new stack without the top (pos)894 llama_grammar_stack next_stack(curr_stack.begin(), curr_stack.end() - 1);895 if (!llama_grammar_is_end_of_sequence(pos + 1)) {896 // if this rule ref is followed by another element, add that to stack897 next_stack.push_back(pos + 1);898 }899 if (!llama_grammar_is_end_of_sequence(subpos)) {900 // if alternate is nonempty, add to stack901 next_stack.push_back(subpos);902 }903 todo.push_back(std::move(next_stack));904 while (!llama_grammar_is_end_of_sequence(subpos)) {905 // scan to end of alternate def906 subpos++;907 }908 if (subpos->type == LLAMA_GRETYPE_ALT) {909 // there's another alternate def of this rule to process910 subpos++;911 } else {912 break;913 }914 } while (true);915 break;916 }917 case LLAMA_GRETYPE_CHAR:918 case LLAMA_GRETYPE_CHAR_NOT:919 case LLAMA_GRETYPE_CHAR_ANY:920 case LLAMA_GRETYPE_TOKEN:921 case LLAMA_GRETYPE_TOKEN_NOT:922 if (std::find(new_stacks.begin(), new_stacks.end(), curr_stack) == new_stacks.end()) {923 // only add the stack if it's not a duplicate of one we already have924 new_stacks.emplace_back(std::move(curr_stack));925 }926 break;927 default:928 // end of alternate (LLAMA_GRETYPE_END, LLAMA_GRETYPE_ALT) or middle of char range929 // (LLAMA_GRETYPE_CHAR_ALT, LLAMA_GRETYPE_CHAR_RNG_UPPER); stack should never be left on930 // those931 GGML_ABORT("fatal error");932 }933 }934}935 936static llama_grammar_candidates llama_grammar_reject_candidates(937 const llama_grammar_rules & rules,938 const llama_grammar_stacks & stacks,939 const llama_grammar_candidates & candidates) {940 GGML_ASSERT(!stacks.empty()); // REVIEW941 942 if (candidates.empty()) {943 return {};944 }945 946 auto rejects = llama_grammar_reject_candidates_for_stack(rules, stacks.front(), candidates);947 948 for (size_t i = 1, size = stacks.size(); i < size; ++i) {949 rejects = llama_grammar_reject_candidates_for_stack(rules, stacks[i], rejects);950 }951 952 return rejects;953}954 955static bool llama_grammar_detect_left_recursion(956 const llama_grammar_rules & rules,957 size_t rule_index,958 std::vector<bool> * rules_visited,959 std::vector<bool> * rules_in_progress,960 std::vector<bool> * rules_may_be_empty) {961 if ((*rules_in_progress)[rule_index]) {962 return true;963 }964 965 (*rules_in_progress)[rule_index] = true;966 967 const llama_grammar_rule & rule = rules[rule_index];968 969 // First check if the rule might produce the empty string. This could be done combined with the second970 // step but it's more readable as two steps.971 bool at_rule_start = true;972 for (size_t i = 0; i < rule.size(); i++) {973 if (llama_grammar_is_end_of_sequence(&rule[i])) {974 if (at_rule_start) {975 (*rules_may_be_empty)[rule_index] = true;976 break;977 }978 at_rule_start = true;979 } else {980 at_rule_start = false;981 }982 }983 984 // Second, recurse into leftmost nonterminals (or next-leftmost as long as the previous nonterminal may985 // be empty)986 bool recurse_into_nonterminal = true;987 for (size_t i = 0; i < rule.size(); i++) {988 if (rule[i].type == LLAMA_GRETYPE_RULE_REF && recurse_into_nonterminal) {989 if (llama_grammar_detect_left_recursion(rules, (size_t)rule[i].value, rules_visited, rules_in_progress, rules_may_be_empty)) {990 return true;991 }992 if (!((*rules_may_be_empty)[(size_t)rule[i].value])) {993 recurse_into_nonterminal = false;994 }995 } else if (llama_grammar_is_end_of_sequence(&rule[i])) {996 recurse_into_nonterminal = true;997 } else {998 recurse_into_nonterminal = false;999 }1000 }1001 1002 (*rules_in_progress)[rule_index] = false;1003 (*rules_visited)[rule_index] = true;1004 1005 return false;1006}1007 1008const llama_grammar_rules & llama_grammar_get_rules(const struct llama_grammar * grammar) {1009 return grammar->rules;1010}1011 1012llama_grammar_stacks & llama_grammar_get_stacks(struct llama_grammar * grammar) {1013 return grammar->stacks;1014}1015 1016static void llama_grammar_accept_chr(1017 struct llama_grammar & grammar,1018 const llama_grammar_stack & stack,1019 uint32_t chr,1020 llama_grammar_stacks & new_stacks) {1021 if (stack.empty()) {1022 return;1023 }1024 1025 const llama_grammar_element * pos = stack.back();1026 1027 // ignore if this turns into a token1028 if (pos->type == LLAMA_GRETYPE_TOKEN || pos->type == LLAMA_GRETYPE_TOKEN_NOT) {1029 return;1030 }1031 1032 auto match = llama_grammar_match_char(pos, chr);1033 if (match.first) {1034 llama_grammar_stack new_stack(stack.begin(), stack.end() - 1);1035 if (!llama_grammar_is_end_of_sequence(match.second)) {1036 new_stack.push_back(match.second);1037 }1038 llama_grammar_advance_stack(grammar.rules, new_stack, new_stacks);1039 }1040}1041 1042void llama_grammar_accept(struct llama_grammar * grammar, uint32_t chr) {1043 llama_grammar_stacks stacks_new;1044 stacks_new.reserve(grammar->stacks.size());1045 1046 for (const auto & stack : grammar->stacks) {1047 llama_grammar_accept_chr(*grammar, stack, chr, stacks_new);1048 }1049 1050 grammar->stacks = std::move(stacks_new);1051}1052 1053llama_grammar_candidates llama_grammar_reject_candidates_for_stack(1054 const llama_grammar_rules & rules,1055 const llama_grammar_stack & stack,1056 const llama_grammar_candidates & candidates) {1057 1058 llama_grammar_candidates rejects;1059 rejects.reserve(candidates.size());1060 1061 if (stack.empty()) {1062 for (const auto & tok : candidates) {1063 if (*tok.code_points != 0 || tok.partial_utf8.n_remain != 0) {1064 rejects.push_back(tok);1065 }1066 }1067 return rejects;1068 }1069 1070 const llama_grammar_element * stack_pos = stack.back();1071 1072 // if the top of the stack is a token rule, then we only need to check the token id1073 if (stack_pos->type == LLAMA_GRETYPE_TOKEN || stack_pos->type == LLAMA_GRETYPE_TOKEN_NOT) {1074 for (const auto & tok : candidates) {1075 if (*tok.code_points == 0) {1076 // reached the end of a token consumed by char rules, reject iff it ended1077 // in a partial response1078 if (tok.partial_utf8.n_remain != 0) {1079 rejects.push_back(tok);1080 }1081 } else if (!llama_grammar_match_token(stack_pos, tok.id)) {1082 rejects.push_back(tok);1083 }1084 }1085 return rejects;1086 }1087 1088 llama_grammar_candidates next_candidates;1089 next_candidates.reserve(candidates.size());1090 1091 for (const auto & tok : candidates) {1092 if (*tok.code_points == 0) {1093 // reached end of full codepoints in token, reject iff it ended in a partial sequence1094 // that cannot satisfy this position in grammar1095 if (tok.partial_utf8.n_remain != 0 &&1096 !llama_grammar_match_partial_char(stack_pos, tok.partial_utf8)) {1097 rejects.push_back(tok);1098 }1099 } else if (llama_grammar_match_char(stack_pos, *tok.code_points).first) {1100 next_candidates.push_back({ tok.index, tok.code_points + 1, tok.partial_utf8, tok.id });1101 } else {1102 rejects.push_back(tok);1103 }1104 }1105 1106 const auto * stack_pos_after = llama_grammar_match_char(stack_pos, 0).second;1107 1108 // update top of stack to next element, if any1109 llama_grammar_stack stack_after(stack.begin(), stack.end() - 1);1110 if (!llama_grammar_is_end_of_sequence(stack_pos_after)) {1111 stack_after.push_back(stack_pos_after);1112 }1113 llama_grammar_stacks next_stacks;1114 llama_grammar_advance_stack(rules, stack_after, next_stacks);1115 1116 auto next_rejects = llama_grammar_reject_candidates(rules, next_stacks, next_candidates);1117 for (const auto & tok : next_rejects) {1118 rejects.push_back({ tok.index, tok.code_points - 1, tok.partial_utf8, tok.id });1119 }1120 1121 return rejects;1122}1123 1124////////////////////1125 1126struct llama_grammar * llama_grammar_init_impl(1127 const struct llama_vocab * vocab,1128 const llama_grammar_element ** rules,1129 size_t n_rules,1130 size_t start_rule_index) {1131 const llama_grammar_element * pos;1132 1133 // copy rule definitions into vectors1134 llama_grammar_rules vec_rules(n_rules);1135 for (size_t i = 0; i < n_rules; i++) {1136 for (pos = rules[i]; pos->type != LLAMA_GRETYPE_END; pos++) {1137 vec_rules[i].push_back(*pos);1138 }1139 vec_rules[i].push_back({LLAMA_GRETYPE_END, 0});1140 }1141 1142 // Check for left recursion1143 std::vector<bool> rules_visited(n_rules);1144 std::vector<bool> rules_in_progress(n_rules);1145 std::vector<bool> rules_may_be_empty(n_rules);1146 for (size_t i = 0; i < n_rules; i++) {1147 if (rules_visited[i]) {1148 continue;1149 }1150 if (llama_grammar_detect_left_recursion(vec_rules, i, &rules_visited, &rules_in_progress, &rules_may_be_empty)) {1151 LLAMA_LOG_ERROR("unsupported grammar, left recursion detected for nonterminal at index %zu", i);1152 return nullptr;1153 }1154 }1155 1156 // loop over alternates of start rule to build initial stacks1157 llama_grammar_stacks stacks;1158 pos = vec_rules[start_rule_index].data();1159 do {1160 llama_grammar_stack stack;1161 if (!llama_grammar_is_end_of_sequence(pos)) {1162 // if alternate is nonempty, add to stack1163 stack.push_back(pos);1164 }1165 llama_grammar_advance_stack(vec_rules, stack, stacks);1166 while (!llama_grammar_is_end_of_sequence(pos)) {1167 // scan to end of alternate def1168 pos++;1169 }1170 if (pos->type == LLAMA_GRETYPE_ALT) {1171 // there's another alternate def of this rule to process1172 pos++;1173 } else {1174 break;1175 }1176 } while (true);1177 1178 // Important: vec_rules has to be moved here, not copied, because stacks contains1179 // pointers to elements of vec_rules. If vec_rules were copied into llama_grammar1180 // then the pointers would be invalidated when the local vec_rules goes out of scope.1181 return new llama_grammar {1182 vocab,1183 std::move(vec_rules),1184 std::move(stacks),1185 /* .partial_utf8 = */ {},1186 /* .lazy = */ false,1187 /* .awaiting_trigger = */ false,1188 /* .trigger_buffer = */ "",1189 /* .trigger_buffer_positions = */ {},1190 /* .trigger_tokens = */ {},1191 /* .trigger_patterns = */ {},1192 };1193}1194 1195struct llama_grammar * llama_grammar_init_impl(1196 const struct llama_vocab * vocab,1197 const char * grammar_str,1198 const char * grammar_root,1199 bool lazy,1200 const char ** trigger_patterns,