mirror of
https://github.com/wassname/prob_jsonformer.git
synced 2026-09-11 12:44:32 +08:00
Merge remote-tracking branch 'posionus/main'
This commit is contained in:
@@ -8,6 +8,8 @@ from prob_jsonformer.logits_processors import (
|
|||||||
StringStoppingCriteria,
|
StringStoppingCriteria,
|
||||||
)
|
)
|
||||||
from prob_jsonformer.choice_tree import choice_tree
|
from prob_jsonformer.choice_tree import choice_tree
|
||||||
|
from prob_jsonformer.type_prefixes import get_prefix_tokens_for_types
|
||||||
|
|
||||||
from termcolor import cprint
|
from termcolor import cprint
|
||||||
from transformers import PreTrainedModel, PreTrainedTokenizer
|
from transformers import PreTrainedModel, PreTrainedTokenizer
|
||||||
import json
|
import json
|
||||||
@@ -37,6 +39,8 @@ class Jsonformer:
|
|||||||
self.json_schema = json_schema
|
self.json_schema = json_schema
|
||||||
self.prompt = prompt
|
self.prompt = prompt
|
||||||
|
|
||||||
|
self.type_prefix_tokens = get_prefix_tokens_for_types(tokenizer)
|
||||||
|
|
||||||
self.number_logit_processor = OutputNumbersTokens(self.tokenizer, self.prompt)
|
self.number_logit_processor = OutputNumbersTokens(self.tokenizer, self.prompt)
|
||||||
self.integer_logit_processor = OutputIntegersTokens(self.tokenizer, self.prompt)
|
self.integer_logit_processor = OutputIntegersTokens(self.tokenizer, self.prompt)
|
||||||
|
|
||||||
@@ -232,6 +236,36 @@ class Jsonformer:
|
|||||||
obj[key] = self.generate_value(schema, obj, key)
|
obj[key] = self.generate_value(schema, obj, key)
|
||||||
return obj
|
return obj
|
||||||
|
|
||||||
|
def choose_type_to_generate(self, possible_types: List[str]) -> str:
|
||||||
|
possible_types = list(set(possible_types)) # remove duplicates
|
||||||
|
self.debug("[choose_type_to_generate]", possible_types)
|
||||||
|
if len(possible_types) < 1:
|
||||||
|
raise ValueError(f"Union type must not be empty")
|
||||||
|
elif len(possible_types) == 1:
|
||||||
|
return possible_types[0]
|
||||||
|
|
||||||
|
prompt = self.get_prompt()
|
||||||
|
input_tensor = self.tokenizer.encode(prompt, return_tensors="pt")
|
||||||
|
output = self.model.forward(input_tensor.to(self.model.device))
|
||||||
|
logits = output.logits[0, -1]
|
||||||
|
|
||||||
|
max_type = None
|
||||||
|
max_logit = -float("inf")
|
||||||
|
for possible_type in possible_types:
|
||||||
|
try:
|
||||||
|
prefix_tokens = self.type_prefix_tokens[possible_type]
|
||||||
|
except KeyError:
|
||||||
|
raise ValueError(f"Unsupported schema type: {possible_type}")
|
||||||
|
max_type_logit = logits[prefix_tokens].max()
|
||||||
|
if max_type_logit > max_logit:
|
||||||
|
max_type = possible_type
|
||||||
|
max_logit = max_type_logit
|
||||||
|
|
||||||
|
if max_type is None:
|
||||||
|
raise Exception("Unable to find best type to generate for union type")
|
||||||
|
self.debug("[choose_type_to_generate]", max_type)
|
||||||
|
return max_type
|
||||||
|
|
||||||
def generate_value(
|
def generate_value(
|
||||||
self,
|
self,
|
||||||
schema: Dict[str, Any],
|
schema: Dict[str, Any],
|
||||||
@@ -239,6 +273,12 @@ class Jsonformer:
|
|||||||
key: Union[str, None] = None,
|
key: Union[str, None] = None,
|
||||||
) -> Any:
|
) -> Any:
|
||||||
schema_type = schema["type"]
|
schema_type = schema["type"]
|
||||||
|
if isinstance(schema_type, list):
|
||||||
|
if key:
|
||||||
|
obj[key] = self.generation_marker
|
||||||
|
else:
|
||||||
|
obj.append(self.generation_marker)
|
||||||
|
schema_type = self.choose_type_to_generate(schema_type)
|
||||||
if schema_type == "number":
|
if schema_type == "number":
|
||||||
if key:
|
if key:
|
||||||
obj[key] = self.generation_marker
|
obj[key] = self.generation_marker
|
||||||
@@ -286,6 +326,8 @@ class Jsonformer:
|
|||||||
else:
|
else:
|
||||||
obj.append(new_obj)
|
obj.append(new_obj)
|
||||||
return self.generate_object(schema["properties"], new_obj)
|
return self.generate_object(schema["properties"], new_obj)
|
||||||
|
elif schema_type == "null":
|
||||||
|
return None
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported schema type: {schema_type}")
|
raise ValueError(f"Unsupported schema type: {schema_type}")
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,32 @@
|
|||||||
|
from transformers import PreTrainedTokenizer
|
||||||
|
from typing import Dict, List
|
||||||
|
import re
|
||||||
|
|
||||||
|
def is_number_prefix(s: str) -> bool:
|
||||||
|
return re.match(r"^[\-\d]+\.?[\d]*$", s)
|
||||||
|
|
||||||
|
def is_boolean_prefix(s: str) -> bool:
|
||||||
|
return 'true'.startswith(s) or 'false'.startswith(s)
|
||||||
|
|
||||||
|
def is_null_prefix(s: str) -> bool:
|
||||||
|
return 'null'.startswith(s)
|
||||||
|
|
||||||
|
def is_string_prefix(s: str) -> bool:
|
||||||
|
return re.match(r'^"[^"]*"?$', s)
|
||||||
|
|
||||||
|
def is_array_prefix(s: str) -> bool:
|
||||||
|
return re.match(r'^\[["\-\d\[{]*$', s)
|
||||||
|
|
||||||
|
def is_object_prefix(s: str) -> bool:
|
||||||
|
return re.match(r'^\{"?$', s)
|
||||||
|
|
||||||
|
def get_prefix_tokens_for_types(tokenizer: PreTrainedTokenizer) -> Dict[str, List[str]]:
|
||||||
|
vocab = tokenizer.vocab.items()
|
||||||
|
return {
|
||||||
|
"number": [v for k, v in vocab if is_number_prefix(k)],
|
||||||
|
"boolean": [v for k, v in vocab if is_boolean_prefix(k)],
|
||||||
|
"null": [v for k, v in vocab if is_null_prefix(k)],
|
||||||
|
"string": [v for k, v in vocab if is_string_prefix(k)],
|
||||||
|
"array": [v for k, v in vocab if is_array_prefix(k)],
|
||||||
|
"object": [v for k, v in vocab if is_object_prefix(k)],
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user