Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 23 additions & 8 deletions src/llama-grammar.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@
#include <cmath>
#include <algorithm>
#include <cstdint>
#include <functional>
#include <iterator>
#include <set>
#include <stdexcept>

Expand Down Expand Up @@ -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<elem_ptr> ptr_less;
std::vector<std::pair<elem_ptr, size_t>> 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()];
}
}
}
Expand Down
43 changes: 43 additions & 0 deletions tests/test-grammar-integration.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<std::string> & passing_strings, const std::vector<std::string> & 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<std::string> & passing_strings, const std::vector<std::string> & failing_strings) {
test(test_desc + ". Schema: " + schema_str, json_schema_to_grammar(json::parse(schema_str), true), passing_strings, failing_strings);
}
Expand Down Expand Up @@ -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;
}