Team Ai
Datasetpublic

echodict/llama.cpp

version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes479downloads
llama-grammar.cpp1511 linesDownload Raw Back to src
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,

Showing the first 1,200 of 1511 lines. Download the file for the rest.