Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
llama-grammar.h165 linesDownload Raw Back to src
1#pragma once2 3#include "llama.h"4 5#include <map>6#include <string>7#include <vector>8 9struct llama_vocab;10 11// grammar element type12enum llama_gretype {13    // end of rule definition14    LLAMA_GRETYPE_END            = 0,15 16    // start of alternate definition for rule17    LLAMA_GRETYPE_ALT            = 1,18 19    // non-terminal element: reference to rule20    LLAMA_GRETYPE_RULE_REF       = 2,21 22    // terminal element: character (code point)23    LLAMA_GRETYPE_CHAR           = 3,24 25    // inverse char(s) ([^a], [^a-b] [^abc])26    LLAMA_GRETYPE_CHAR_NOT       = 4,27 28    // modifies a preceding LLAMA_GRETYPE_CHAR or LLAMA_GRETYPE_CHAR_ALT to29    // be an inclusive range ([a-z])30    LLAMA_GRETYPE_CHAR_RNG_UPPER = 5,31 32    // modifies a preceding LLAMA_GRETYPE_CHAR or33    // LLAMA_GRETYPE_CHAR_RNG_UPPER to add an alternate char to match ([ab], [a-zA])34    LLAMA_GRETYPE_CHAR_ALT       = 6,35 36    // any character (.)37    LLAMA_GRETYPE_CHAR_ANY       = 7,38};39 40typedef struct llama_grammar_element {41    enum llama_gretype type;42    uint32_t           value; // Unicode code point or rule ID43} llama_grammar_element;44 45struct llama_partial_utf8 {46    uint32_t value;    // bit value so far (unshifted)47    int      n_remain; // num bytes remaining; -1 indicates invalid sequence48};49 50struct llama_grammar_candidate {51    size_t               index;52    const uint32_t     * code_points;53    llama_partial_utf8   partial_utf8;54};55 56using llama_grammar_rule  = std::vector<      llama_grammar_element>;57using llama_grammar_stack = std::vector<const llama_grammar_element *>;58 59using llama_grammar_rules      = std::vector<llama_grammar_rule>;60using llama_grammar_stacks     = std::vector<llama_grammar_stack>;61using llama_grammar_candidates = std::vector<llama_grammar_candidate>;62 63// TODO: remove, needed for tests atm64const llama_grammar_rules  & llama_grammar_get_rules (const struct llama_grammar * grammar);65      llama_grammar_stacks & llama_grammar_get_stacks(      struct llama_grammar * grammar);66 67// takes a set of possible pushdown stacks on a grammar, which are required to68// be positioned at a character range (see `llama_grammar_advance_stack`), and69// produces the N possible stacks if the given char is accepted at those70// positions71void llama_grammar_accept(struct llama_grammar * grammar, uint32_t chr);72 73std::vector<llama_grammar_candidate> llama_grammar_reject_candidates_for_stack(74        const llama_grammar_rules      & rules,75        const llama_grammar_stack      & stack,76        const llama_grammar_candidates & candidates);77 78struct llama_grammar_parser {79    std::map<std::string, uint32_t> symbol_ids;80 81    llama_grammar_rules rules;82 83    llama_grammar_stack c_rules() const;84 85    uint32_t get_symbol_id(const char * src, size_t len);86    uint32_t generate_symbol_id(const std::string & base_name);87 88    void add_rule(uint32_t rule_id, const llama_grammar_rule & rule);89 90    const char * parse_alternates(91            const char        * src,92            const std::string & rule_name,93            uint32_t            rule_id,94            bool                is_nested);95 96    const char * parse_sequence(97            const char         * src,98            const std::string  & rule_name,99            llama_grammar_rule & rule,100            bool               is_nested);101 102    const char * parse_rule(const char * src);103 104    bool parse(const char * src);105    void print(FILE * file);106};107 108struct llama_grammar {109    // note: allow null vocab for testing (not great)110    const llama_vocab * vocab;111 112    const llama_grammar_rules  rules;  // TODO: shared ptr113          llama_grammar_stacks stacks;114 115    // buffer for partially generated UTF-8 sequence from accepted tokens116    llama_partial_utf8 partial_utf8;117 118    // lazy grammars wait for trigger words or tokens before constraining the sampling.119    // we still ahve trigger_tokens for non-lazy grammars to force printing of special trigger tokens.120    // (useful e.g. for tool_choice=required)121    bool                     lazy             = false;122    bool                     awaiting_trigger = false; // Initialized to true for lazy grammars only123    std::string              trigger_buffer;           // Output buffered by lazy grammar. Will be cleared once trigger is found.124    std::vector<llama_token> trigger_tokens;           // Tokens that trigger a lazy grammar, or tokens to force printing of (even if special).125    std::vector<std::string> trigger_words;126};127 128//129// internal API130//131 132// note: needed for tests (not great)133struct llama_grammar * llama_grammar_init_impl(134        const struct llama_vocab * vocab,135        const llama_grammar_element ** rules,136        size_t n_rules,137        size_t start_rule_index);138 139struct llama_grammar * llama_grammar_init_impl(140        const struct llama_vocab * vocab,141                      const char * grammar_str,142                      const char * grammar_root,143                              bool lazy,144                     const char ** trigger_words,145                            size_t num_trigger_words,146               const llama_token * trigger_tokens,147                            size_t num_trigger_tokens);148 149void llama_grammar_free_impl(struct llama_grammar * grammar);150 151struct llama_grammar * llama_grammar_clone_impl(const struct llama_grammar & grammar);152 153// TODO: move the API below as member functions of llama_grammar154void llama_grammar_apply_impl(155        const struct llama_grammar & grammar,156            llama_token_data_array * cur_p);157 158void llama_grammar_accept_impl(159              struct llama_grammar & grammar,160                       llama_token   token);161 162void llama_grammar_accept_str(163              struct llama_grammar & grammar,164                 const std::string & piece);165