Team Ai
Datasetpublic

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.

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes3.1kdownloads
llama-grammar.cpp1525 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            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,

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

Brunobkr/llama.cpp_AlgMor24_github · Team Ai