From ce1277549012a33e5c2360f42bf53aaf1b95e528 Mon Sep 17 00:00:00 2001 From: Andrei Betlen Date: Tue, 6 Feb 2024 18:50:56 -0500 Subject: [PATCH 01/13] Update llama.cpp --- vendor/llama.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/vendor/llama.cpp b/vendor/llama.cpp index b08f22c882..213d1439fa 160000 --- a/vendor/llama.cpp +++ b/vendor/llama.cpp @@ -1 +1 @@ -Subproject commit b08f22c882a1443e6b97081f3ce718a4d1a741f8 +Subproject commit 213d1439fadefe182f69c5f7e8dd3b4b6572ebcb From 2ef7ba3aed572609fbf7292adb125e41e5279a15 Mon Sep 17 00:00:00 2001 From: Andrei Betlen Date: Thu, 8 Feb 2024 01:07:44 -0500 Subject: [PATCH 02/13] misc: rename grammar test --- tests/{test_grammar.py => test_llama_grammar.py} | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename tests/{test_grammar.py => test_llama_grammar.py} (100%) diff --git a/tests/test_grammar.py b/tests/test_llama_grammar.py similarity index 100% rename from tests/test_grammar.py rename to tests/test_llama_grammar.py From b5fca911b57a23565c55c31802fb9603a0c6497c Mon Sep 17 00:00:00 2001 From: Andrei Betlen Date: Thu, 8 Feb 2024 01:08:18 -0500 Subject: [PATCH 03/13] feat: Move tokenizer to own module --- llama_cpp/llama.py | 69 ++------------------------ llama_cpp/llama_tokenizer.py | 96 ++++++++++++++++++++++++++++++++++++ 2 files changed, 100 insertions(+), 65 deletions(-) create mode 100644 llama_cpp/llama_tokenizer.py diff --git a/llama_cpp/llama.py b/llama_cpp/llama.py index bad75dfaf7..30ae3b56cd 100644 --- a/llama_cpp/llama.py +++ b/llama_cpp/llama.py @@ -2,7 +2,6 @@ import os import sys -import abc import uuid import time import multiprocessing @@ -15,7 +14,6 @@ Iterator, Deque, Callable, - Any, ) from collections import deque @@ -31,6 +29,10 @@ LlamaDiskCache, # type: ignore LlamaRAMCache, # type: ignore ) +from .llama_tokenizer import ( + BaseLlamaTokenizer, + LlamaTokenizer +) import llama_cpp.llama_cpp as llama_cpp import llama_cpp.llama_chat_format as llama_chat_format @@ -1747,69 +1749,6 @@ def longest_token_prefix(a: Sequence[int], b: Sequence[int]): return longest_prefix -class BaseLlamaTokenizer(abc.ABC): - @abc.abstractmethod - def tokenize(self, text: bytes, add_bos: bool = True, special: bool = True) -> List[int]: - raise NotImplementedError - - @abc.abstractmethod - def detokenize(self, tokens: List[int], prev_tokens: Optional[List[int]] = None) -> bytes: - raise NotImplementedError - - -class LlamaTokenizer(BaseLlamaTokenizer): - def __init__(self, llama: Llama): - self.llama = llama - self._model = llama._model # type: ignore - - def tokenize(self, text: bytes, add_bos: bool = True, special: bool = True) -> List[int]: - return self._model.tokenize(text, add_bos=add_bos, special=special) - - def detokenize(self, tokens: List[int], prev_tokens: Optional[List[int]] = None) -> bytes: - if prev_tokens is not None: - return self._model.detokenize(tokens[len(prev_tokens):]) - else: - return self._model.detokenize(tokens) - - def encode(self, text: str, add_bos: bool = True, special: bool = True) -> List[int]: - return self.tokenize( - text.encode("utf-8", errors="ignore"), add_bos=add_bos, special=special - ) - - def decode(self, tokens: List[int]) -> str: - return self.detokenize(tokens).decode("utf-8", errors="ignore") - - @classmethod - def from_ggml_file(cls, path: str) -> "LlamaTokenizer": - return cls(Llama(model_path=path, vocab_only=True)) - - -class LlamaHFTokenizer(BaseLlamaTokenizer): - def __init__(self, hf_tokenizer: Any): - self.hf_tokenizer = hf_tokenizer - - def tokenize(self, text: bytes, add_bos: bool = True, special: bool = True) -> List[int]: - return self.hf_tokenizer.encode(text.decode("utf-8", errors="ignore"), add_special_tokens=special) - - def detokenize(self, tokens: List[int], prev_tokens: Optional[List[int]] = None) -> bytes: - if prev_tokens is not None: - text = self.hf_tokenizer.decode(tokens).encode("utf-8", errors="ignore") - prev_text = self.hf_tokenizer.decode(prev_tokens).encode("utf-8", errors="ignore") - return text[len(prev_text):] - else: - return self.hf_tokenizer.decode(tokens).encode("utf-8", errors="ignore") - - @classmethod - def from_pretrained(cls, pretrained_model_name_or_path: str) -> "LlamaHFTokenizer": - try: - from transformers import AutoTokenizer - except ImportError: - raise ImportError( - "The `transformers` library is required to use the `HFTokenizer`." - "You can install it with `pip install transformers`." - ) - hf_tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path=pretrained_model_name_or_path) - return cls(hf_tokenizer) class LlamaState: diff --git a/llama_cpp/llama_tokenizer.py b/llama_cpp/llama_tokenizer.py new file mode 100644 index 0000000000..0ad3c3afee --- /dev/null +++ b/llama_cpp/llama_tokenizer.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +import abc +from typing import ( + List, + Optional, + Any, +) + +import llama_cpp +from llama_cpp.llama_types import List + + +class BaseLlamaTokenizer(abc.ABC): + @abc.abstractmethod + def tokenize( + self, text: bytes, add_bos: bool = True, special: bool = True + ) -> List[int]: + raise NotImplementedError + + @abc.abstractmethod + def detokenize( + self, tokens: List[int], prev_tokens: Optional[List[int]] = None + ) -> bytes: + raise NotImplementedError + + +class LlamaTokenizer(BaseLlamaTokenizer): + def __init__(self, llama: llama_cpp.Llama): + self.llama = llama + self._model = llama._model # type: ignore + + def tokenize( + self, text: bytes, add_bos: bool = True, special: bool = True + ) -> List[int]: + return self._model.tokenize(text, add_bos=add_bos, special=special) + + def detokenize( + self, tokens: List[int], prev_tokens: Optional[List[int]] = None + ) -> bytes: + if prev_tokens is not None: + return self._model.detokenize(tokens[len(prev_tokens) :]) + else: + return self._model.detokenize(tokens) + + def encode( + self, text: str, add_bos: bool = True, special: bool = True + ) -> List[int]: + return self.tokenize( + text.encode("utf-8", errors="ignore"), add_bos=add_bos, special=special + ) + + def decode(self, tokens: List[int]) -> str: + return self.detokenize(tokens).decode("utf-8", errors="ignore") + + @classmethod + def from_ggml_file(cls, path: str) -> "LlamaTokenizer": + return cls(llama_cpp.Llama(model_path=path, vocab_only=True)) + + +class LlamaHFTokenizer(BaseLlamaTokenizer): + def __init__(self, hf_tokenizer: Any): + self.hf_tokenizer = hf_tokenizer + + def tokenize( + self, text: bytes, add_bos: bool = True, special: bool = True + ) -> List[int]: + return self.hf_tokenizer.encode( + text.decode("utf-8", errors="ignore"), add_special_tokens=special + ) + + def detokenize( + self, tokens: List[int], prev_tokens: Optional[List[int]] = None + ) -> bytes: + if prev_tokens is not None: + text = self.hf_tokenizer.decode(tokens).encode("utf-8", errors="ignore") + prev_text = self.hf_tokenizer.decode(prev_tokens).encode( + "utf-8", errors="ignore" + ) + return text[len(prev_text) :] + else: + return self.hf_tokenizer.decode(tokens).encode("utf-8", errors="ignore") + + @classmethod + def from_pretrained(cls, pretrained_model_name_or_path: str) -> "LlamaHFTokenizer": + try: + from transformers import AutoTokenizer + except ImportError: + raise ImportError( + "The `transformers` library is required to use the `HFTokenizer`." + "You can install it with `pip install transformers`." + ) + hf_tokenizer = AutoTokenizer.from_pretrained( + pretrained_model_name_or_path=pretrained_model_name_or_path + ) + return cls(hf_tokenizer) From 85d3374b4d5892e51e27b9973f9ce3623e076e2a Mon Sep 17 00:00:00 2001 From: Andrei Betlen Date: Thu, 8 Feb 2024 01:13:28 -0500 Subject: [PATCH 04/13] fix: broken import --- llama_cpp/server/model.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/llama_cpp/server/model.py b/llama_cpp/server/model.py index 6d8ec24670..5308dc2a89 100644 --- a/llama_cpp/server/model.py +++ b/llama_cpp/server/model.py @@ -6,6 +6,7 @@ import llama_cpp import llama_cpp.llama_speculative as llama_speculative +import llama_cpp.llama_tokenizer as llama_tokenizer from llama_cpp.server.settings import ModelSettings @@ -95,7 +96,7 @@ def load_llama_from_model_settings(settings: ModelSettings) -> llama_cpp.Llama: tokenizer: Optional[llama_cpp.BaseLlamaTokenizer] = None if settings.hf_pretrained_model_name_or_path is not None: - tokenizer = llama_cpp.LlamaHFTokenizer.from_pretrained(settings.hf_pretrained_model_name_or_path) + tokenizer = llama_tokenizer.LlamaHFTokenizer.from_pretrained(settings.hf_pretrained_model_name_or_path) draft_model = None if settings.draft_model is not None: From dfc1b173414b550f8f5be1b94430af16b53a63cb Mon Sep 17 00:00:00 2001 From: Andrei Betlen Date: Thu, 8 Feb 2024 23:38:12 -0500 Subject: [PATCH 05/13] Update llama.cpp --- vendor/llama.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/vendor/llama.cpp b/vendor/llama.cpp index 213d1439fa..8e6a9d2de0 160000 --- a/vendor/llama.cpp +++ b/vendor/llama.cpp @@ -1 +1 @@ -Subproject commit 213d1439fadefe182f69c5f7e8dd3b4b6572ebcb +Subproject commit 8e6a9d2de0096af7120606c74ee2f26684e87b41 From e16f06e6eb555947f4404c20732921c8ea76c4f7 Mon Sep 17 00:00:00 2001 From: Andrei Betlen Date: Fri, 9 Feb 2024 02:02:13 -0500 Subject: [PATCH 06/13] fix: revert _create_completions. --- llama_cpp/llama.py | 20 ++++++++++++-------- 1 file changed, 12 insertions(+), 8 deletions(-) diff --git a/llama_cpp/llama.py b/llama_cpp/llama.py index bad75dfaf7..f445fb0e9a 100644 --- a/llama_cpp/llama.py +++ b/llama_cpp/llama.py @@ -948,8 +948,7 @@ def logit_bias_processor( if stream: remaining_tokens = completion_tokens[returned_tokens:] - prev_tokens = completion_tokens[:returned_tokens] - remaining_text = self.detokenize(completion_tokens, prev_tokens) + remaining_text = self.detokenize(remaining_tokens) remaining_length = len(remaining_text) # We want to avoid yielding any characters from @@ -971,13 +970,13 @@ def logit_bias_processor( for token in remaining_tokens: if token == self.token_bos(): continue - token_end_position += len(remaining_text) + token_end_position += len(self.detokenize([token])) # Check if stop sequence is in the token if token_end_position > ( remaining_length - first_stop_position ): break - token_str = remaining_text.decode( + token_str = self.detokenize([token]).decode( "utf-8", errors="ignore" ) text_offset = len(prompt) + len( @@ -1002,7 +1001,11 @@ def logit_bias_processor( } top_logprob.update({token_str: current_logprobs[int(token)]}) logprobs_or_none = { - "tokens": [token_str], + "tokens": [ + self.detokenize([token]).decode( + "utf-8", errors="ignore" + ) + ], "text_offset": [text_offset], "token_logprobs": [current_logprobs[int(token)]], "top_logprobs": [top_logprob], @@ -1015,7 +1018,9 @@ def logit_bias_processor( "model": model_name, "choices": [ { - "text": token_str, + "text": self.detokenize([token]).decode( + "utf-8", errors="ignore" + ), "index": 0, "logprobs": logprobs_or_none, "finish_reason": None, @@ -1027,7 +1032,7 @@ def logit_bias_processor( decode_success = False for i in range(1, len(remaining_tokens) + 1): try: - bs = remaining_text + bs = self.detokenize(remaining_tokens[:i]) ts = bs.decode("utf-8") decode_success = True break @@ -1063,7 +1068,6 @@ def logit_bias_processor( if len(completion_tokens) >= max_tokens: text = self.detokenize(completion_tokens) - finish_reason = "length" break From 63b0c37836169baa71c04484e5344294928bd359 Mon Sep 17 00:00:00 2001 From: Andrei Betlen Date: Fri, 9 Feb 2024 13:36:58 -0500 Subject: [PATCH 07/13] Update llama.cpp --- vendor/llama.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/vendor/llama.cpp b/vendor/llama.cpp index 8e6a9d2de0..4b7b38bef5 160000 --- a/vendor/llama.cpp +++ b/vendor/llama.cpp @@ -1 +1 @@ -Subproject commit 8e6a9d2de0096af7120606c74ee2f26684e87b41 +Subproject commit 4b7b38bef5addbd31f453871d79647fbae6bec8a From 837f39d9f3dd0531809b2ba6782840a020412ede Mon Sep 17 00:00:00 2001 From: David Brown Date: Fri, 9 Feb 2024 20:09:21 +0100 Subject: [PATCH 08/13] feat: implement required attributes in json_schema_to_gbnf --- llama_cpp/llama_grammar.py | 56 ++++++++++++++++++++-------------- tests/test_grammar.py | 61 +++++++++++++++++++++++++++++++++++--- 2 files changed, 91 insertions(+), 26 deletions(-) diff --git a/llama_cpp/llama_grammar.py b/llama_cpp/llama_grammar.py index d8ef563c2d..bc72a94926 100644 --- a/llama_cpp/llama_grammar.py +++ b/llama_cpp/llama_grammar.py @@ -1392,14 +1392,14 @@ def print_grammar(file: TextIO, state: parse_state) -> None: SPACE_RULE = '" "?' PRIMITIVE_RULES = { - "boolean": '("true" | "false") space', - "number": '("-"? ([0-9] | [1-9] [0-9]*)) ("." [0-9]+)? ([eE] [-+]? [0-9]+)? space', - "integer": '("-"? ([0-9] | [1-9] [0-9]*)) space', + "boolean": '("true" | "false")', + "number": '("-"? ([0-9] | [1-9] [0-9]*)) ("." [0-9]+)? ([eE] [-+]? [0-9]+)?', + "integer": '("-"? ([0-9] | [1-9] [0-9]*))', "string": r""" "\"" ( [^"\\] | "\\" (["\\/bfnrt] | "u" [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F]) - )* "\"" space """, - "null": '"null" space', + )* "\"" """, + "null": '"null"', } INVALID_RULE_CHARS_RE = re.compile(r"[^a-zA-Z0-9-]+") @@ -1419,19 +1419,29 @@ def _format_literal(self, literal: str): ) return f'"{escaped}"' - def _add_rule(self, name: str, rule: str): + def _add_rule( + self, name: str, rule: str, is_required: bool = True, with_space: bool = False + ): esc_name = INVALID_RULE_CHARS_RE.sub("-", name) - if esc_name not in self._rules or self._rules[esc_name] == rule: + if not is_required: + esc_name += "-or-null" + + complete_rule = rule if is_required else f"({rule} | {PRIMITIVE_RULES['null']})" + if with_space: + complete_rule += " space" + + if esc_name not in self._rules or self._rules[esc_name] == complete_rule: key = esc_name else: i = 0 while f"{esc_name}{i}" in self._rules: i += 1 key = f"{esc_name}{i}" - self._rules[key] = rule + + self._rules[key] = complete_rule return key - def visit(self, schema: Dict[str, Any], name: str) -> str: + def visit(self, schema: Dict[str, Any], name: str, is_required: bool = True) -> str: rule_name = name or "root" if "$defs" in schema: @@ -1448,14 +1458,16 @@ def visit(self, schema: Dict[str, Any], name: str) -> str: ) ) ) - return self._add_rule(rule_name, rule) + return self._add_rule(rule_name, rule, is_required, False) elif "const" in schema: - return self._add_rule(rule_name, self._format_literal(schema["const"])) + return self._add_rule( + rule_name, self._format_literal(schema["const"]), is_required, False + ) elif "enum" in schema: rule = " | ".join((self._format_literal(v) for v in schema["enum"])) - return self._add_rule(rule_name, rule) + return self._add_rule(rule_name, rule, is_required, False) elif "$ref" in schema: ref = schema["$ref"] @@ -1465,12 +1477,10 @@ def visit(self, schema: Dict[str, Any], name: str) -> str: def_schema = self._defs[def_name] return self.visit(def_schema, f'{name}{"-" if name else ""}{def_name}') - - schema_type: Optional[str] = schema.get("type") # type: ignore + schema_type: Optional[str] = schema.get("type") # type: ignore assert isinstance(schema_type, str), f"Unrecognized schema: {schema}" if schema_type == "object" and "properties" in schema: - # TODO: `required` keyword prop_order = self._prop_order prop_pairs = sorted( schema["properties"].items(), @@ -1481,30 +1491,32 @@ def visit(self, schema: Dict[str, Any], name: str) -> str: rule = '"{" space' for i, (prop_name, prop_schema) in enumerate(prop_pairs): prop_rule_name = self.visit( - prop_schema, f'{name}{"-" if name else ""}{prop_name}' + prop_schema, + f'{name}{"-" if name else ""}{prop_name}', + "required" not in schema or prop_name in schema["required"], ) if i > 0: rule += ' "," space' rule += rf' {self._format_literal(prop_name)} space ":" space {prop_rule_name}' - rule += ' "}" space' + rule += ' "}"' - return self._add_rule(rule_name, rule) + return self._add_rule(rule_name, rule, is_required, True) elif schema_type == "array" and "items" in schema: # TODO `prefixItems` keyword item_rule_name = self.visit( schema["items"], f'{name}{"-" if name else ""}item' ) - rule = ( - f'"[" space ({item_rule_name} ("," space {item_rule_name})*)? "]" space' - ) - return self._add_rule(rule_name, rule) + rule = f'"[" space ({item_rule_name} ("," space {item_rule_name})*)? "]"' + return self._add_rule(rule_name, rule, is_required, True) else: assert schema_type in PRIMITIVE_RULES, f"Unrecognized schema: {schema}" return self._add_rule( "root" if rule_name == "root" else schema_type, PRIMITIVE_RULES[schema_type], + is_required, + True, ) def format_grammar(self): diff --git a/tests/test_grammar.py b/tests/test_grammar.py index cb221880a6..7ac2c66699 100644 --- a/tests/test_grammar.py +++ b/tests/test_grammar.py @@ -42,7 +42,7 @@ class B(BaseModel): "a": {"$ref": "#/$defs/A"}, "b": {"title": "B", "type": "integer"}, }, - "required": ["a", "b"], + "required": ["a"], "title": "B", "type": "object", } @@ -51,9 +51,18 @@ class B(BaseModel): assert grammar.grammar is not None + assert ( + llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None) + == r"""space ::= " "? +integer ::= ("-"? ([0-9] | [1-9] [0-9]*)) space +a-A ::= "{" space "\"a\"" space ":" space integer "}" space +integer-or-null ::= (("-"? ([0-9] | [1-9] [0-9]*)) | "null") space +root ::= "{" space "\"a\"" space ":" space a-A "," space "\"b\"" space ":" space integer-or-null "}" space""" + ) + def test_grammar_anyof(): - sch = { + schema = { "properties": { "temperature": { "description": "The temperature mentioned", @@ -73,6 +82,50 @@ def test_grammar_anyof(): "type": "object", } - grammar = llama_cpp.LlamaGrammar.from_json_schema(json.dumps(sch)) + grammar = llama_cpp.LlamaGrammar.from_json_schema(json.dumps(schema)) + + assert grammar.grammar is not None + + assert ( + llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None) + == r"""space ::= " "? +number ::= ("-"? ([0-9] | [1-9] [0-9]*)) ("." [0-9]+)? ([eE] [-+]? [0-9]+)? space +unit-0 ::= "\"celsius\"" | "\"fahrenheit\"" +null ::= "null" space +unit ::= unit-0 | null +root ::= "{" space "\"temperature\"" space ":" space number "," space "\"unit\"" space ":" space unit "}" space""" + ) + + +def test_grammar_nested_object(): + schema = { + "type": "object", + "properties": { + "test": {"type": "string"}, + "nested": { + "type": "object", + "properties": {"other": {"type": "string"}}, + "required": [], + }, + }, + "required": ["test"], + } + + grammar = llama_cpp.LlamaGrammar.from_json_schema(json.dumps(schema)) + + assert grammar.grammar is not None - assert grammar.grammar is not None \ No newline at end of file + assert ( + llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None) + == r"""space ::= " "? +string-or-null ::= ( "\"" ( + [^"\\] | + "\\" (["\\/bfnrt] | "u" [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F]) + )* "\"" | "null") space +nested-or-null ::= ("{" space "\"other\"" space ":" space string-or-null "}" | "null") space +string ::= "\"" ( + [^"\\] | + "\\" (["\\/bfnrt] | "u" [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F]) + )* "\"" space +root ::= "{" space "\"nested\"" space ":" space nested-or-null "," space "\"test\"" space ":" space string "}" space""" + ) From 9ab518d72c948fd0fd335aaf6ba7651228d5e313 Mon Sep 17 00:00:00 2001 From: David Brown Date: Mon, 12 Feb 2024 08:53:37 +0100 Subject: [PATCH 09/13] add treat_optional_as_nullable option --- llama_cpp/llama_grammar.py | 23 +++++++++++++++++------ tests/test_llama_grammar.py | 23 +++++++++++++++++++++-- 2 files changed, 38 insertions(+), 8 deletions(-) diff --git a/llama_cpp/llama_grammar.py b/llama_cpp/llama_grammar.py index bc72a94926..7ba6730d20 100644 --- a/llama_cpp/llama_grammar.py +++ b/llama_cpp/llama_grammar.py @@ -1408,10 +1408,11 @@ def print_grammar(file: TextIO, state: parse_state) -> None: class SchemaConverter: - def __init__(self, prop_order): + def __init__(self, prop_order, treat_optional_as_nullable: bool = False): self._prop_order = prop_order self._rules = {"space": SPACE_RULE} self._defs: Dict[str, Any] = {} + self._treat_optional_as_nullable = treat_optional_as_nullable def _format_literal(self, literal: str): escaped: str = GRAMMAR_LITERAL_ESCAPE_RE.sub( @@ -1423,10 +1424,12 @@ def _add_rule( self, name: str, rule: str, is_required: bool = True, with_space: bool = False ): esc_name = INVALID_RULE_CHARS_RE.sub("-", name) - if not is_required: + complete_rule = rule + + if self._treat_optional_as_nullable and not is_required: esc_name += "-or-null" + complete_rule = f"({complete_rule} | {PRIMITIVE_RULES['null']})" - complete_rule = rule if is_required else f"({rule} | {PRIMITIVE_RULES['null']})" if with_space: complete_rule += " space" @@ -1481,6 +1484,7 @@ def visit(self, schema: Dict[str, Any], name: str, is_required: bool = True) -> assert isinstance(schema_type, str), f"Unrecognized schema: {schema}" if schema_type == "object" and "properties" in schema: + # TODO: `required` keyword when treat_optional_as_nullable=False prop_order = self._prop_order prop_pairs = sorted( schema["properties"].items(), @@ -1490,10 +1494,13 @@ def visit(self, schema: Dict[str, Any], name: str, is_required: bool = True) -> rule = '"{" space' for i, (prop_name, prop_schema) in enumerate(prop_pairs): + is_prop_required = ( + "required" not in schema or prop_name in schema["required"] + ) prop_rule_name = self.visit( prop_schema, f'{name}{"-" if name else ""}{prop_name}', - "required" not in schema or prop_name in schema["required"], + is_prop_required, ) if i > 0: rule += ' "," space' @@ -1523,10 +1530,14 @@ def format_grammar(self): return "\n".join((f"{name} ::= {rule}" for name, rule in self._rules.items())) -def json_schema_to_gbnf(schema: str, prop_order: Optional[List[str]] = None): +def json_schema_to_gbnf( + schema: str, + prop_order: Optional[List[str]] = None, + treat_optional_as_nullable: bool = False, +): prop_order = prop_order or [] schema = json.loads(schema) prop_order = {name: idx for idx, name in enumerate(prop_order)} - converter = SchemaConverter(prop_order) + converter = SchemaConverter(prop_order, treat_optional_as_nullable) converter.visit(schema, "") return converter.format_grammar() diff --git a/tests/test_llama_grammar.py b/tests/test_llama_grammar.py index 7ac2c66699..7ab8c4bb3c 100644 --- a/tests/test_llama_grammar.py +++ b/tests/test_llama_grammar.py @@ -52,7 +52,15 @@ class B(BaseModel): assert grammar.grammar is not None assert ( - llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None) + llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None, False) + == r"""space ::= " "? +integer ::= ("-"? ([0-9] | [1-9] [0-9]*)) space +a-A ::= "{" space "\"a\"" space ":" space integer "}" space +root ::= "{" space "\"a\"" space ":" space a-A "," space "\"b\"" space ":" space integer "}" space""" + ) + + assert ( + llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None, True) == r"""space ::= " "? integer ::= ("-"? ([0-9] | [1-9] [0-9]*)) space a-A ::= "{" space "\"a\"" space ":" space integer "}" space @@ -116,7 +124,18 @@ def test_grammar_nested_object(): assert grammar.grammar is not None assert ( - llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None) + llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None, False) + == r"""space ::= " "? +string ::= "\"" ( + [^"\\] | + "\\" (["\\/bfnrt] | "u" [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F]) + )* "\"" space +nested ::= "{" space "\"other\"" space ":" space string "}" space +root ::= "{" space "\"nested\"" space ":" space nested "," space "\"test\"" space ":" space string "}" space""" + ) + + assert ( + llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None, True) == r"""space ::= " "? string-or-null ::= ( "\"" ( [^"\\] | From 34658905b9504593e7e570540169252a96819d57 Mon Sep 17 00:00:00 2001 From: David Brown Date: Mon, 12 Feb 2024 15:37:15 +0100 Subject: [PATCH 10/13] implement optional attributes --- llama_cpp/llama_grammar.py | 29 +++++++++++++++++++++++------ tests/test_llama_grammar.py | 4 ++-- 2 files changed, 25 insertions(+), 8 deletions(-) diff --git a/llama_cpp/llama_grammar.py b/llama_cpp/llama_grammar.py index 7ba6730d20..e97729bd23 100644 --- a/llama_cpp/llama_grammar.py +++ b/llama_cpp/llama_grammar.py @@ -1484,7 +1484,6 @@ def visit(self, schema: Dict[str, Any], name: str, is_required: bool = True) -> assert isinstance(schema_type, str), f"Unrecognized schema: {schema}" if schema_type == "object" and "properties" in schema: - # TODO: `required` keyword when treat_optional_as_nullable=False prop_order = self._prop_order prop_pairs = sorted( schema["properties"].items(), @@ -1492,7 +1491,8 @@ def visit(self, schema: Dict[str, Any], name: str, is_required: bool = True) -> key=lambda kv: (prop_order.get(kv[0], len(prop_order)), kv[0]), ) - rule = '"{" space' + rule = "" + previous_is_prop_required = None for i, (prop_name, prop_schema) in enumerate(prop_pairs): is_prop_required = ( "required" not in schema or prop_name in schema["required"] @@ -1502,10 +1502,27 @@ def visit(self, schema: Dict[str, Any], name: str, is_required: bool = True) -> f'{name}{"-" if name else ""}{prop_name}', is_prop_required, ) - if i > 0: - rule += ' "," space' - rule += rf' {self._format_literal(prop_name)} space ":" space {prop_rule_name}' - rule += ' "}"' + prop_rule = rf'{self._format_literal(prop_name)} space ":" space {prop_rule_name}' + if i == 0: + rule += prop_rule + previous_is_prop_required = is_prop_required + else: + if self._treat_optional_as_nullable or ( + previous_is_prop_required and is_prop_required + ): + rule = f'{rule} "," space {prop_rule}' + previous_is_prop_required = True + elif previous_is_prop_required and not is_prop_required: + rule = f'{rule} ("," space {prop_rule})?' + previous_is_prop_required = True + elif not previous_is_prop_required and is_prop_required: + rule = f'({rule} "," space)? {prop_rule}' + previous_is_prop_required = True + elif not previous_is_prop_required and not is_prop_required: + rule = f'({rule} | {prop_rule} | {rule} "," space {prop_rule})' + previous_is_prop_required = False + + rule = '"{" space ' + rule + ' "}"' return self._add_rule(rule_name, rule, is_required, True) diff --git a/tests/test_llama_grammar.py b/tests/test_llama_grammar.py index 7ab8c4bb3c..43c888940c 100644 --- a/tests/test_llama_grammar.py +++ b/tests/test_llama_grammar.py @@ -56,7 +56,7 @@ class B(BaseModel): == r"""space ::= " "? integer ::= ("-"? ([0-9] | [1-9] [0-9]*)) space a-A ::= "{" space "\"a\"" space ":" space integer "}" space -root ::= "{" space "\"a\"" space ":" space a-A "," space "\"b\"" space ":" space integer "}" space""" +root ::= "{" space "\"a\"" space ":" space a-A ("," space "\"b\"" space ":" space integer)? "}" space""" ) assert ( @@ -131,7 +131,7 @@ def test_grammar_nested_object(): "\\" (["\\/bfnrt] | "u" [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F]) )* "\"" space nested ::= "{" space "\"other\"" space ":" space string "}" space -root ::= "{" space "\"nested\"" space ":" space nested "," space "\"test\"" space ":" space string "}" space""" +root ::= "{" space ("\"nested\"" space ":" space nested "," space)? "\"test\"" space ":" space string "}" space""" ) assert ( From d06bbf010bdd051894328657aee89cc183e04dc9 Mon Sep 17 00:00:00 2001 From: David Brown Date: Mon, 12 Feb 2024 16:03:25 +0100 Subject: [PATCH 11/13] add space after last value in objects --- llama_cpp/llama_grammar.py | 2 +- tests/test_llama_grammar.py | 18 +++++++++--------- 2 files changed, 10 insertions(+), 10 deletions(-) diff --git a/llama_cpp/llama_grammar.py b/llama_cpp/llama_grammar.py index e97729bd23..cfc0901a12 100644 --- a/llama_cpp/llama_grammar.py +++ b/llama_cpp/llama_grammar.py @@ -1522,7 +1522,7 @@ def visit(self, schema: Dict[str, Any], name: str, is_required: bool = True) -> rule = f'({rule} | {prop_rule} | {rule} "," space {prop_rule})' previous_is_prop_required = False - rule = '"{" space ' + rule + ' "}"' + rule = '"{" space ' + rule + ' space "}"' return self._add_rule(rule_name, rule, is_required, True) diff --git a/tests/test_llama_grammar.py b/tests/test_llama_grammar.py index 43c888940c..2bbbf789c0 100644 --- a/tests/test_llama_grammar.py +++ b/tests/test_llama_grammar.py @@ -55,17 +55,17 @@ class B(BaseModel): llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None, False) == r"""space ::= " "? integer ::= ("-"? ([0-9] | [1-9] [0-9]*)) space -a-A ::= "{" space "\"a\"" space ":" space integer "}" space -root ::= "{" space "\"a\"" space ":" space a-A ("," space "\"b\"" space ":" space integer)? "}" space""" +a-A ::= "{" space "\"a\"" space ":" space integer space "}" space +root ::= "{" space "\"a\"" space ":" space a-A ("," space "\"b\"" space ":" space integer)? space "}" space""" ) assert ( llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None, True) == r"""space ::= " "? integer ::= ("-"? ([0-9] | [1-9] [0-9]*)) space -a-A ::= "{" space "\"a\"" space ":" space integer "}" space +a-A ::= "{" space "\"a\"" space ":" space integer space "}" space integer-or-null ::= (("-"? ([0-9] | [1-9] [0-9]*)) | "null") space -root ::= "{" space "\"a\"" space ":" space a-A "," space "\"b\"" space ":" space integer-or-null "}" space""" +root ::= "{" space "\"a\"" space ":" space a-A "," space "\"b\"" space ":" space integer-or-null space "}" space""" ) @@ -101,7 +101,7 @@ def test_grammar_anyof(): unit-0 ::= "\"celsius\"" | "\"fahrenheit\"" null ::= "null" space unit ::= unit-0 | null -root ::= "{" space "\"temperature\"" space ":" space number "," space "\"unit\"" space ":" space unit "}" space""" +root ::= "{" space "\"temperature\"" space ":" space number "," space "\"unit\"" space ":" space unit space "}" space""" ) @@ -130,8 +130,8 @@ def test_grammar_nested_object(): [^"\\] | "\\" (["\\/bfnrt] | "u" [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F]) )* "\"" space -nested ::= "{" space "\"other\"" space ":" space string "}" space -root ::= "{" space ("\"nested\"" space ":" space nested "," space)? "\"test\"" space ":" space string "}" space""" +nested ::= "{" space "\"other\"" space ":" space string space "}" space +root ::= "{" space ("\"nested\"" space ":" space nested "," space)? "\"test\"" space ":" space string space "}" space""" ) assert ( @@ -141,10 +141,10 @@ def test_grammar_nested_object(): [^"\\] | "\\" (["\\/bfnrt] | "u" [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F]) )* "\"" | "null") space -nested-or-null ::= ("{" space "\"other\"" space ":" space string-or-null "}" | "null") space +nested-or-null ::= ("{" space "\"other\"" space ":" space string-or-null space "}" | "null") space string ::= "\"" ( [^"\\] | "\\" (["\\/bfnrt] | "u" [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F] [0-9a-fA-F]) )* "\"" space -root ::= "{" space "\"nested\"" space ":" space nested-or-null "," space "\"test\"" space ":" space string "}" space""" +root ::= "{" space "\"nested\"" space ":" space nested-or-null "," space "\"test\"" space ":" space string space "}" space""" ) From 5acc92c9e94c08096242e33e71ba17c9ddf99230 Mon Sep 17 00:00:00 2001 From: David Brown Date: Mon, 12 Feb 2024 16:20:30 +0100 Subject: [PATCH 12/13] expose treat_optional_as_nullable option in from_json_schema --- llama_cpp/llama_grammar.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/llama_cpp/llama_grammar.py b/llama_cpp/llama_grammar.py index cfc0901a12..cdc31a5684 100644 --- a/llama_cpp/llama_grammar.py +++ b/llama_cpp/llama_grammar.py @@ -81,9 +81,15 @@ def from_json_schema( cls, json_schema: str, verbose: bool = True, + treat_optional_as_nullable: bool = False, ) -> "LlamaGrammar": """Convert a JSON schema to a Llama grammar.""" - return cls.from_string(json_schema_to_gbnf(json_schema), verbose=verbose) + return cls.from_string( + json_schema_to_gbnf( + json_schema, treat_optional_as_nullable=treat_optional_as_nullable + ), + verbose=verbose, + ) @classmethod def from_file(cls, file: Union[str, Path], verbose: bool = True) -> "LlamaGrammar": From 6d62834663c715d4fa4fd8060621c35b94d2e3ab Mon Sep 17 00:00:00 2001 From: David Brown Date: Mon, 12 Feb 2024 16:31:19 +0100 Subject: [PATCH 13/13] small refacto --- llama_cpp/llama_grammar.py | 9 +++------ tests/test_llama_grammar.py | 16 ++++++++++++---- 2 files changed, 15 insertions(+), 10 deletions(-) diff --git a/llama_cpp/llama_grammar.py b/llama_cpp/llama_grammar.py index cdc31a5684..6f5e7d258f 100644 --- a/llama_cpp/llama_grammar.py +++ b/llama_cpp/llama_grammar.py @@ -1498,7 +1498,7 @@ def visit(self, schema: Dict[str, Any], name: str, is_required: bool = True) -> ) rule = "" - previous_is_prop_required = None + previous_is_prop_required = False for i, (prop_name, prop_schema) in enumerate(prop_pairs): is_prop_required = ( "required" not in schema or prop_name in schema["required"] @@ -1511,22 +1511,19 @@ def visit(self, schema: Dict[str, Any], name: str, is_required: bool = True) -> prop_rule = rf'{self._format_literal(prop_name)} space ":" space {prop_rule_name}' if i == 0: rule += prop_rule - previous_is_prop_required = is_prop_required else: if self._treat_optional_as_nullable or ( previous_is_prop_required and is_prop_required ): rule = f'{rule} "," space {prop_rule}' - previous_is_prop_required = True elif previous_is_prop_required and not is_prop_required: rule = f'{rule} ("," space {prop_rule})?' - previous_is_prop_required = True elif not previous_is_prop_required and is_prop_required: rule = f'({rule} "," space)? {prop_rule}' - previous_is_prop_required = True elif not previous_is_prop_required and not is_prop_required: rule = f'({rule} | {prop_rule} | {rule} "," space {prop_rule})' - previous_is_prop_required = False + + previous_is_prop_required |= is_prop_required rule = '"{" space ' + rule + ' space "}"' diff --git a/tests/test_llama_grammar.py b/tests/test_llama_grammar.py index 2bbbf789c0..459f05e151 100644 --- a/tests/test_llama_grammar.py +++ b/tests/test_llama_grammar.py @@ -52,7 +52,9 @@ class B(BaseModel): assert grammar.grammar is not None assert ( - llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None, False) + llama_cpp.llama_grammar.json_schema_to_gbnf( + json.dumps(schema), treat_optional_as_nullable=False + ) == r"""space ::= " "? integer ::= ("-"? ([0-9] | [1-9] [0-9]*)) space a-A ::= "{" space "\"a\"" space ":" space integer space "}" space @@ -60,7 +62,9 @@ class B(BaseModel): ) assert ( - llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None, True) + llama_cpp.llama_grammar.json_schema_to_gbnf( + json.dumps(schema), treat_optional_as_nullable=True + ) == r"""space ::= " "? integer ::= ("-"? ([0-9] | [1-9] [0-9]*)) space a-A ::= "{" space "\"a\"" space ":" space integer space "}" space @@ -124,7 +128,9 @@ def test_grammar_nested_object(): assert grammar.grammar is not None assert ( - llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None, False) + llama_cpp.llama_grammar.json_schema_to_gbnf( + json.dumps(schema), treat_optional_as_nullable=False + ) == r"""space ::= " "? string ::= "\"" ( [^"\\] | @@ -135,7 +141,9 @@ def test_grammar_nested_object(): ) assert ( - llama_cpp.llama_grammar.json_schema_to_gbnf(json.dumps(schema), None, True) + llama_cpp.llama_grammar.json_schema_to_gbnf( + json.dumps(schema), treat_optional_as_nullable=True + ) == r"""space ::= " "? string-or-null ::= ( "\"" ( [^"\\] |