diff --git a/src/llama-grammar.cpp b/src/llama-grammar.cpp index c685346b630..8c823820b2b 100644 --- a/src/llama-grammar.cpp +++ b/src/llama-grammar.cpp @@ -7,6 +7,8 @@ #include #include #include +#include +#include #include #include @@ -1335,14 +1337,27 @@ struct llama_grammar * llama_grammar_clone_impl(const struct llama_grammar & gra }; // redirect elements in stacks to point to new rules - for (size_t is = 0; is < result->stacks.size(); is++) { - for (size_t ie = 0; ie < result->stacks[is].size(); ie++) { - for (size_t ir0 = 0; ir0 < grammar.rules.size(); ir0++) { - for (size_t ir1 = 0; ir1 < grammar.rules[ir0].size(); ir1++) { - if (grammar.stacks[is][ie] == &grammar.rules[ir0][ir1]) { - result->stacks[is][ie] = &result->rules[ir0][ir1]; - } - } + // find the rule of each element with a binary search on the rule start addresses: a search of all elements is too slow for large grammars (the server clones the grammar on each speculative step) + using elem_ptr = const llama_grammar_element *; + const std::less ptr_less; + std::vector> starts; + starts.reserve(grammar.rules.size()); + for (size_t ir = 0; ir < grammar.rules.size(); ir++) { + if (!grammar.rules[ir].empty()) { + starts.emplace_back(grammar.rules[ir].data(), ir); + } + } + std::sort(starts.begin(), starts.end(), [&](const auto & a, const auto & b) { return ptr_less(a.first, b.first); }); + for (auto & stack : result->stacks) { + for (auto & pe : stack) { + auto it = std::upper_bound(starts.begin(), starts.end(), pe, [&](elem_ptr v, const auto & e) { return ptr_less(v, e.first); }); + if (it == starts.begin()) { + continue; + } + const size_t ir = std::prev(it)->second; + const auto & rule = grammar.rules[ir]; + if (ptr_less(pe, rule.data() + rule.size())) { + pe = &result->rules[ir][pe - rule.data()]; } } } diff --git a/tests/test-grammar-integration.cpp b/tests/test-grammar-integration.cpp index eb4b7c78f50..1574ab5b087 100644 --- a/tests/test-grammar-integration.cpp +++ b/tests/test-grammar-integration.cpp @@ -188,6 +188,48 @@ static void test(const std::string & test_desc, const std::string & grammar_str, static void test_grammar(const std::string & test_desc, const std::string & grammar_str, const std::vector & passing_strings, const std::vector & failing_strings) { test(test_desc + ". Grammar: " + grammar_str, grammar_str, passing_strings, failing_strings); } + +static bool stacks_point_into_rules(llama_grammar * grammar) { + const auto & rules = llama_grammar_get_rules(grammar); + for (const auto & stack : llama_grammar_get_stacks(grammar)) { + for (const llama_grammar_element * pe : stack) { + bool found = false; + for (const auto & rule : rules) { + found = found || (!rule.empty() && pe >= rule.data() && pe < rule.data() + rule.size()); + } + if (!found) { + return false; + } + } + } + return true; +} + +static void test_clone() { + fprintf(stderr, "⚫ Testing grammar clone at each position of an input\n"); + const std::string grammar_str = json_schema_to_grammar(json::parse(R"""({ + "type": "object", + "properties": { + "name": {"type": "string"}, + "tags": {"type": "array", "items": {"type": "string"}}, + "n": {"type": "integer"} + }, + "required": ["name", "n"] + })"""), true); + const std::string input = R"""({"name": "a", "n": 42, "tags": ["x", "y"]})"""; + for (size_t k = 0; k <= input.size(); k++) { + auto * grammar = build_grammar(grammar_str); + for (const auto & in : parse_tokens(input.substr(0, k))) { + llama_grammar_accept_token(*grammar, in.token, in.piece); + } + auto * clone = llama_grammar_clone_impl(*grammar); + llama_grammar_free_impl(grammar); + assert(stacks_point_into_rules(clone)); + assert(match_string(input.substr(k), clone)); + llama_grammar_free_impl(clone); + } + fprintf(stdout, " ✅︎\n"); +} static void test_schema(const std::string & test_desc, const std::string & schema_str, const std::vector & passing_strings, const std::vector & failing_strings) { test(test_desc + ". Schema: " + schema_str, json_schema_to_grammar(json::parse(schema_str), true), passing_strings, failing_strings); } @@ -1488,6 +1530,7 @@ int main() { test_failure_missing_root_symbol(); test_custom_root_symbol_check(); test_json_schema(); + test_clone(); fprintf(stdout, "All tests passed.\n"); return 0; }