mirror of
https://github.com/wassname/prob_jsonformer.git
synced 2026-09-09 11:29:57 +08:00
update
This commit is contained in:
+3
-1
@@ -1 +1,3 @@
|
||||
jsonllm/__pycache__
|
||||
jsonllm/__pycache__
|
||||
.venv
|
||||
workspace.ipynb
|
||||
|
||||
+217
@@ -0,0 +1,217 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 1,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"/home/ubuntu/jsonllm/.venv/lib/python3.10/site-packages/tqdm/auto.py:21: TqdmWarning: IProgress not found. Please update jupyter and ipywidgets. See https://ipywidgets.readthedocs.io/en/stable/user_install.html\n",
|
||||
" from .autonotebook import tqdm as notebook_tqdm\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Loading model and tokenizer...\n",
|
||||
"Loaded model and tokenizer\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from transformers import AutoModelForCausalLM, AutoTokenizer\n",
|
||||
"\n",
|
||||
"print(\"Loading model and tokenizer...\")\n",
|
||||
"model_name = \"databricks/dolly-v2-12b\"\n",
|
||||
"model = AutoModelForCausalLM.from_pretrained(model_name, use_cache=True)\n",
|
||||
"tokenizer = AutoTokenizer.from_pretrained(model_name, use_fast=True, use_cache=True)\n",
|
||||
"print(\"Loaded model and tokenizer\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Generating...\n",
|
||||
"{\n",
|
||||
" temperature: \u001b[32m2.2225\u001b[0m,\n",
|
||||
" humidity: \u001b[32m1.0\u001b[0m,\n",
|
||||
" wind_speed: {\n",
|
||||
" value: \u001b[32m0.0\u001b[0m,\n",
|
||||
" unit: \u001b[32m\"value\"\u001b[0m\n",
|
||||
" }\n",
|
||||
"}\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from jsonformer.format import highlight_values\n",
|
||||
"from jsonformer.main import Jsonformer\n",
|
||||
"\n",
|
||||
"weather_schema = {\n",
|
||||
" \"type\": \"object\",\n",
|
||||
" \"properties\": {\n",
|
||||
" \"temperature\": {\"type\": \"number\"},\n",
|
||||
" \"humidity\": {\n",
|
||||
" \"type\": \"number\",\n",
|
||||
" },\n",
|
||||
" \"wind_speed\": {\n",
|
||||
" \"type\": \"object\",\n",
|
||||
" \"properties\": {\n",
|
||||
" \"value\": {\"type\": \"number\"},\n",
|
||||
" \"unit\": {\"type\": \"string\"},\n",
|
||||
" },\n",
|
||||
" },\n",
|
||||
" },\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
"builder = Jsonformer(\n",
|
||||
" model=model,\n",
|
||||
" tokenizer=tokenizer,\n",
|
||||
" json_schema=weather_schema,\n",
|
||||
" prompt=\"generate the weather\",\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"print(\"Generating...\")\n",
|
||||
"output = builder()\n",
|
||||
"\n",
|
||||
"highlight_values(output)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 8,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Generating...\n",
|
||||
"{\n",
|
||||
" make: \u001b[32m\"Ford\"\u001b[0m,\n",
|
||||
" model: \u001b[32m\"Mustang\"\u001b[0m,\n",
|
||||
" year: \u001b[32m10.0\u001b[0m,\n",
|
||||
" colors: [\n",
|
||||
" \u001b[32m\"red\"\u001b[0m,\n",
|
||||
" \u001b[32m\"white\"\u001b[0m,\n",
|
||||
" \u001b[32m\"blue\"\u001b[0m,\n",
|
||||
" \u001b[32m\"black\"\u001b[0m,\n",
|
||||
" \u001b[32m\"yellow\"\u001b[0m,\n",
|
||||
" \u001b[32m\"orange\"\u001b[0m,\n",
|
||||
" \u001b[32m\"green\"\u001b[0m,\n",
|
||||
" \u001b[32m\"pink\"\u001b[0m,\n",
|
||||
" \u001b[32m\"purple\"\u001b[0m,\n",
|
||||
" \u001b[32m\"violet\"\u001b[0m\n",
|
||||
" ]\n",
|
||||
"}\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"car = {\n",
|
||||
" \"type\": \"object\",\n",
|
||||
" \"properties\": {\n",
|
||||
" \"make\": {\"type\": \"string\"},\n",
|
||||
" \"model\": {\"type\": \"string\"},\n",
|
||||
" \"year\": {\"type\": \"number\"},\n",
|
||||
" \"colors\": {\n",
|
||||
" \"type\": \"array\",\n",
|
||||
" \"items\": {\"type\": \"string\"},\n",
|
||||
" }\n",
|
||||
" },\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
"builder = Jsonformer(\n",
|
||||
" model=model,\n",
|
||||
" tokenizer=tokenizer,\n",
|
||||
" json_schema=car,\n",
|
||||
" prompt=\"generate an example car\",\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"print(\"Generating...\")\n",
|
||||
"output = builder()\n",
|
||||
"\n",
|
||||
"highlight_values(output)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 9,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"{\n",
|
||||
" temperature: \u001b[32m24\u001b[0m,\n",
|
||||
" humidity: \u001b[32m48.0\u001b[0m,\n",
|
||||
" wind_speed: {\n",
|
||||
" value: \u001b[32m12\u001b[0m,\n",
|
||||
" unit: \u001b[32m\"mph\"\u001b[0m\n",
|
||||
" }\n",
|
||||
"}\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"test = {\n",
|
||||
" \"temperature\": 24,\n",
|
||||
" \"humidity\": 48.0,\n",
|
||||
" \"wind_speed\": {\n",
|
||||
" \"value\": 12,\n",
|
||||
" \"unit\": \"mph\"\n",
|
||||
" }\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
"highlight_values(test)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": ".venv",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.10.11"
|
||||
},
|
||||
"orig_nbformat": 4
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 2
|
||||
}
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 60 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 42 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 81 KiB |
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,47 @@
|
||||
from jsonformer.format import highlight_values
|
||||
from jsonformer.main import Jsonformer
|
||||
from transformers import (
|
||||
AutoModelForCausalLM,
|
||||
AutoTokenizer,
|
||||
)
|
||||
import torch
|
||||
|
||||
|
||||
weather_schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"temperature": {"type": "number"},
|
||||
"humidity": {
|
||||
"type": "number",
|
||||
},
|
||||
"wind_speed": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"value": {"type": "number"},
|
||||
"unit": {"type": "string"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
print("Loading model and tokenizer...")
|
||||
model_name = "databricks/dolly-v2-12b"
|
||||
model = AutoModelForCausalLM.from_pretrained(model_name, use_cache=True)
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_name, use_fast=True, use_cache=True)
|
||||
print("Loaded model and tokenizer")
|
||||
|
||||
|
||||
builder = Jsonformer(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
json_schema=weather_schema,
|
||||
prompt="generate the weather",
|
||||
debug=True,
|
||||
)
|
||||
|
||||
print("Generating...")
|
||||
output = builder()
|
||||
|
||||
highlight_values(
|
||||
output,
|
||||
)
|
||||
@@ -0,0 +1,24 @@
|
||||
from termcolor import colored
|
||||
|
||||
|
||||
def highlight_values(value):
|
||||
def recursive_print(obj, indent=0, is_last_element=True):
|
||||
if isinstance(obj, dict):
|
||||
print("{")
|
||||
last_key = list(obj.keys())[-1]
|
||||
for key, value in obj.items():
|
||||
print(f"{' ' * (indent + 2)}{key}: ", end="")
|
||||
recursive_print(value, indent + 2, key == last_key)
|
||||
print(f"{' ' * indent}}}", end=",\n" if not is_last_element else "\n")
|
||||
elif isinstance(obj, list):
|
||||
print("[")
|
||||
for index, value in enumerate(obj):
|
||||
print(f"{' ' * (indent + 2)}", end="")
|
||||
recursive_print(value, indent + 2, index == len(obj) - 1)
|
||||
print(f"{' ' * indent}]", end=",\n" if not is_last_element else "\n")
|
||||
else:
|
||||
if isinstance(obj, str):
|
||||
obj = f'"{obj}"'
|
||||
print(colored(obj, "green"), end=",\n" if not is_last_element else "\n")
|
||||
|
||||
recursive_print(value)
|
||||
@@ -3,10 +3,6 @@ import torch
|
||||
|
||||
|
||||
class NumberStoppingCriteria(StoppingCriteria):
|
||||
"""
|
||||
This class can be used to stop generation when there is a repeated decimal point in the generated text.
|
||||
"""
|
||||
|
||||
def __init__(self, tokenizer: PreTrainedTokenizer, precision: int = 2):
|
||||
self.tokenizer = tokenizer
|
||||
self.precision = precision
|
||||
@@ -17,16 +13,11 @@ class NumberStoppingCriteria(StoppingCriteria):
|
||||
scores: torch.FloatTensor,
|
||||
) -> bool:
|
||||
decoded = self.tokenizer.decode(input_ids[0], skip_special_tokens=True)
|
||||
if ".." in decoded:
|
||||
print("Stopping because of ..")
|
||||
return True
|
||||
|
||||
if decoded.strip().count(".") > 1:
|
||||
print("Stopping because of multiple .")
|
||||
if decoded.count(".") > 1:
|
||||
return True
|
||||
|
||||
if (
|
||||
decoded.strip().count(".") == 1
|
||||
decoded.count(".") == 1
|
||||
and len(decoded.strip().split(".")[1]) > self.precision
|
||||
):
|
||||
return True
|
||||
@@ -36,7 +27,9 @@ class NumberStoppingCriteria(StoppingCriteria):
|
||||
|
||||
class OutputNumbersTokens(LogitsWarper):
|
||||
def __init__(self, tokenizer: PreTrainedTokenizer, prompt: str):
|
||||
self.whitelist_tokens = [tokenizer.eos_token_id]
|
||||
self.whitelist_tokens = [
|
||||
# tokenizer.eos_token_id
|
||||
]
|
||||
self.tokenized_prompt = tokenizer(prompt, return_tensors="pt")
|
||||
|
||||
for token_str, token_id in tokenizer.get_vocab().items():
|
||||
@@ -0,0 +1,192 @@
|
||||
from typing import List, Union, Dict, Any
|
||||
|
||||
from jsonformer.logits_processors import NumberStoppingCriteria, OutputNumbersTokens
|
||||
from transformers import PreTrainedModel, PreTrainedTokenizer
|
||||
import json
|
||||
|
||||
GENERATION_MARKER = "|GENERATION|"
|
||||
|
||||
|
||||
class Jsonformer:
|
||||
value: Dict[str, Any] = {}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: PreTrainedModel,
|
||||
tokenizer: PreTrainedTokenizer,
|
||||
json_schema: Dict[str, Any],
|
||||
prompt: str,
|
||||
debug: bool = False,
|
||||
max_array_length: int = 10,
|
||||
):
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
self.json_schema = json_schema
|
||||
self.prompt = prompt
|
||||
|
||||
self.number_logit_processor = OutputNumbersTokens(self.tokenizer, self.prompt)
|
||||
self.number_stop_criteria = NumberStoppingCriteria(self.tokenizer, 3)
|
||||
|
||||
self.generation_marker = "|GENERATION|"
|
||||
self.debug_on = debug
|
||||
self.max_array_length = max_array_length
|
||||
|
||||
def debug(self, *args, **kwargs):
|
||||
if self.debug_on:
|
||||
print(*args, **kwargs)
|
||||
|
||||
def generate_number(self) -> float:
|
||||
prompt = self.get_prompt()
|
||||
self.debug("[generate_number] prompt", prompt)
|
||||
response = self.model.generate(
|
||||
self.tokenizer.encode(prompt, return_tensors="pt"),
|
||||
max_new_tokens=6,
|
||||
num_return_sequences=1,
|
||||
logits_processor=[self.number_logit_processor],
|
||||
stopping_criteria=[self.number_stop_criteria],
|
||||
temperature=1.2,
|
||||
pad_token_id=self.tokenizer.eos_token_id,
|
||||
)
|
||||
response = self.tokenizer.decode(response[0], skip_special_tokens=True)
|
||||
self.debug("[generate_number] response", response)
|
||||
response = response[len(prompt) :]
|
||||
response = response.strip().rstrip(".")
|
||||
|
||||
try:
|
||||
return float(response)
|
||||
except ValueError:
|
||||
print("ValueError")
|
||||
return
|
||||
|
||||
def generate_boolean(self) -> bool:
|
||||
prompt = self.get_prompt()
|
||||
self.debug("[generate_boolean] prompt", prompt)
|
||||
|
||||
input_tensor = self.tokenizer.encode(prompt, return_tensors="pt")
|
||||
output = self.model.forward(input_tensor)
|
||||
logits = output.logits[0, -1]
|
||||
|
||||
true_token_id = self.tokenizer.convert_tokens_to_ids("true")
|
||||
false_token_id = self.tokenizer.convert_tokens_to_ids("false")
|
||||
|
||||
true_logits = logits[true_token_id]
|
||||
false_logits = logits[false_token_id]
|
||||
|
||||
if true_logits > false_logits:
|
||||
return True
|
||||
elif false_logits > true_logits:
|
||||
return False
|
||||
else:
|
||||
print("Failed to generate a valid boolean value")
|
||||
return None
|
||||
|
||||
def generate_string(self) -> str:
|
||||
prompt = self.get_prompt()
|
||||
self.debug("[generate_string] prompt", prompt)
|
||||
response = self.model.generate(
|
||||
self.tokenizer.encode(prompt, return_tensors="pt"),
|
||||
max_new_tokens=8,
|
||||
num_return_sequences=1,
|
||||
temperature=1.3,
|
||||
pad_token_id=self.tokenizer.eos_token_id,
|
||||
)
|
||||
response = self.tokenizer.decode(response[0], skip_special_tokens=True)
|
||||
|
||||
response = response[len(prompt) :].strip()
|
||||
|
||||
self.debug("[generate_string] response", response)
|
||||
split = response.split('"')
|
||||
assert len(split) >= 2
|
||||
return split[1]
|
||||
|
||||
def generate_object(
|
||||
self, properties: Dict[str, Any], obj: Dict[str, Any]
|
||||
) -> Dict[str, Any]:
|
||||
# self.debug("[generate_object] properties", properties)
|
||||
for key, schema in properties.items():
|
||||
obj[key] = self.generate_value(schema, obj, key)
|
||||
return obj
|
||||
|
||||
def generate_value(
|
||||
self,
|
||||
schema: Dict[str, Any],
|
||||
obj: Union[Dict[str, Any], List[Any]],
|
||||
key: str | None = None,
|
||||
) -> Any:
|
||||
schema_type = schema["type"]
|
||||
if schema_type == "number":
|
||||
obj[key if key else -1] = self.generation_marker
|
||||
return self.generate_number()
|
||||
elif schema_type == "boolean":
|
||||
obj[key if key else -1] = self.generation_marker
|
||||
return self.generate_boolean()
|
||||
elif schema_type == "string":
|
||||
obj[key if key else -1] = self.generation_marker
|
||||
return self.generate_string()
|
||||
elif schema_type == "array":
|
||||
new_array = []
|
||||
obj[key] = new_array
|
||||
return self.generate_array(schema["items"], new_array)
|
||||
elif schema_type == "object":
|
||||
new_obj = {}
|
||||
obj[key if key else -1] = new_obj
|
||||
return self.generate_object(schema["properties"], new_obj)
|
||||
else:
|
||||
raise ValueError(f"Unsupported schema type: {schema_type}")
|
||||
|
||||
def generate_array(self, item_schema: Dict[str, Any], obj: Dict[str, Any]) -> list:
|
||||
for i in range(self.max_array_length):
|
||||
element = self.generate_value(item_schema, obj)
|
||||
obj[-1] = element
|
||||
|
||||
obj.append(self.generation_marker)
|
||||
input_prompt = self.get_prompt()
|
||||
obj.pop()
|
||||
input_tensor = self.tokenizer.encode(input_prompt, return_tensors="pt")
|
||||
output = self.model.forward(input_tensor)
|
||||
logits = output.logits[0, -1]
|
||||
|
||||
close_bracket_token_id = self.tokenizer.convert_tokens_to_ids("]")
|
||||
comma_token_id = self.tokenizer.convert_tokens_to_ids(", ")
|
||||
self.debug(
|
||||
"[generate_array] token_ids", close_bracket_token_id, comma_token_id
|
||||
)
|
||||
close_bracket_logits = logits[close_bracket_token_id]
|
||||
comma_logits = logits[comma_token_id]
|
||||
|
||||
self.debug(
|
||||
"[generate_array] close_bracket_logits",
|
||||
close_bracket_logits,
|
||||
"comma_logits",
|
||||
comma_logits,
|
||||
)
|
||||
|
||||
if close_bracket_logits > comma_logits:
|
||||
break
|
||||
|
||||
return obj
|
||||
|
||||
def get_prompt(self):
|
||||
template = """{prompt}\nOutput result in the following JSON schema format:\n{schema}\nResult: {progress}"""
|
||||
progress = json.dumps(self.value)
|
||||
gen_marker_index = progress.find(f'"{self.generation_marker}"')
|
||||
if gen_marker_index != -1:
|
||||
progress = progress[:gen_marker_index]
|
||||
else:
|
||||
print("Failed to find generation marker")
|
||||
|
||||
prompt = template.format(
|
||||
prompt=self.prompt,
|
||||
schema=json.dumps(self.json_schema),
|
||||
progress=progress,
|
||||
)
|
||||
|
||||
return prompt
|
||||
|
||||
def __call__(self) -> Dict[str, Any]:
|
||||
self.value = {}
|
||||
|
||||
generated_data = self.generate_object(
|
||||
self.json_schema["properties"], self.value
|
||||
)
|
||||
return generated_data
|
||||
-173
@@ -1,173 +0,0 @@
|
||||
import json
|
||||
import random
|
||||
from jsonllm.logits_processors import NumberStoppingCriteria, OutputNumbersTokens
|
||||
from transformers import (
|
||||
PreTrainedTokenizer,
|
||||
PreTrainedModel,
|
||||
AutoModelForCausalLM,
|
||||
AutoTokenizer,
|
||||
)
|
||||
from typing import Any, Dict
|
||||
|
||||
|
||||
class JSONLLM:
|
||||
value: Dict[str, Any] = {}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: PreTrainedModel,
|
||||
tokenizer: PreTrainedTokenizer,
|
||||
json_schema: Dict[str, Any],
|
||||
prompt: str,
|
||||
):
|
||||
self.model = model
|
||||
self.tokenizer = tokenizer
|
||||
self.json_schema = json_schema
|
||||
self.prompt = prompt
|
||||
|
||||
self.number_logit_processor = OutputNumbersTokens(self.tokenizer, self.prompt)
|
||||
self.number_stop_criteria = NumberStoppingCriteria(self.tokenizer)
|
||||
|
||||
# def generate(
|
||||
# self, text: str, max_length: int = 100, forced_bos_token_id: list | None = None
|
||||
# ) -> str:
|
||||
# # print prompt in red
|
||||
# print("generate")
|
||||
# print("\033[91m {}\033[00m".format(text))
|
||||
# input_ids = self.tokenizer.encode(text, return_tensors="pt")
|
||||
# output = self.model.generate(
|
||||
# input_ids,
|
||||
# max_length=max_length,
|
||||
# num_return_sequences=1,
|
||||
# forced_bos_token_id=forced_bos_token_id,
|
||||
# )
|
||||
# decoded_output = self.tokenizer.decode(output[0], skip_special_tokens=True)
|
||||
# return decoded_output.strip()
|
||||
|
||||
def generate_number(self, suffix="") -> float:
|
||||
print("generate_number", suffix)
|
||||
prompt = self.get_prompt() + suffix
|
||||
|
||||
# print prompt in red
|
||||
# number = model.generate(
|
||||
# tokenizer.encode(prompt, return_tensors="pt"),
|
||||
# max_new_tokens=5,
|
||||
# num_return_sequences=1,
|
||||
# logits_processor=[a],
|
||||
# stopping_criteria=[NumberStoppingCriteria(tokenizer)],
|
||||
# )
|
||||
|
||||
print("\033[91m {}\033[00m".format(prompt))
|
||||
response = self.model.generate(
|
||||
self.tokenizer.encode(prompt, return_tensors="pt"),
|
||||
max_new_tokens=5,
|
||||
num_return_sequences=1,
|
||||
logits_processor=[self.number_logit_processor],
|
||||
stopping_criteria=[self.number_stop_criteria],
|
||||
)
|
||||
|
||||
response = self.tokenizer.decode(response[0], skip_special_tokens=True)
|
||||
|
||||
print("\033[94m {}\033[00m".format(response))
|
||||
try:
|
||||
return float(response)
|
||||
except ValueError:
|
||||
print("ValueError")
|
||||
return
|
||||
|
||||
def generate_boolean(self, suffix="") -> bool:
|
||||
prompt = self.get_prompt()
|
||||
true_token_id = self.tokenizer.encode("true", add_special_tokens=False)[0]
|
||||
false_token_id = self.tokenizer.encode("false", add_special_tokens=False)[0]
|
||||
|
||||
response = self.generate(
|
||||
prompt, forced_bos_token_id=[true_token_id, false_token_id]
|
||||
).lower()
|
||||
|
||||
if response == "true":
|
||||
return True
|
||||
else:
|
||||
return False
|
||||
|
||||
def generate_array(
|
||||
self, item_schema: Dict[str, Any], obj: Dict[str, Any], suffix=""
|
||||
) -> list:
|
||||
array_length = random.randint(0, 5)
|
||||
return [self.generate_value(item_schema, obj) for _ in range(array_length)]
|
||||
|
||||
# add stopping criteria with "
|
||||
def generate_string(self) -> str:
|
||||
prompt = self.get_prompt()
|
||||
response = self.generate(prompt)
|
||||
return response
|
||||
|
||||
def generate_object(
|
||||
self, properties: Dict[str, Any], obj: Dict[str, Any], suffix=""
|
||||
) -> Dict[str, Any]:
|
||||
print("generate_object", properties)
|
||||
|
||||
for key, schema in properties.items():
|
||||
value = self.generate_value(schema, obj, suffix=f'"{key}": ')
|
||||
|
||||
obj[key] = value
|
||||
return obj
|
||||
|
||||
def generate_value(self, schema: Dict[str, Any], obj: Dict[str, Any], suffix=""):
|
||||
schema_type = schema["type"]
|
||||
if schema_type == "number":
|
||||
return self.generate_number(suffix=suffix)
|
||||
elif schema_type == "boolean":
|
||||
return self.generate_boolean(suffix=suffix)
|
||||
elif schema_type == "array":
|
||||
return self.generate_array(schema["items"], obj, suffix=suffix)
|
||||
elif schema_type == "object":
|
||||
return self.generate_object(schema["properties"], obj, suffix=suffix)
|
||||
else:
|
||||
raise ValueError(f"Unsupported schema type: {schema_type}")
|
||||
|
||||
def get_prompt(self):
|
||||
template = """{prompt}\nMake sure to output in the following format:\n{schema}\n {progress}"""
|
||||
progress = json.dumps(self.value)
|
||||
|
||||
progress = progress.rstrip("}").rstrip("]").rstrip(",")
|
||||
|
||||
prompt = template.format(
|
||||
prompt=self.prompt,
|
||||
schema=json.dumps(self.json_schema),
|
||||
progress=progress,
|
||||
)
|
||||
|
||||
return prompt
|
||||
|
||||
def __call__(self) -> Dict[str, Any]:
|
||||
self.value = {}
|
||||
generated_data = self.generate_object(
|
||||
self.json_schema["properties"], self.value
|
||||
)
|
||||
return generated_data
|
||||
|
||||
|
||||
model_name = "EleutherAI/gpt-neo-1.3B"
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
model_name,
|
||||
)
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
||||
|
||||
weather_schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"temperature": {"type": "number"},
|
||||
# "humidity": {"type": "number"},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
jsonllm = JSONLLM(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
json_schema=weather_schema,
|
||||
prompt="Generate a weather object",
|
||||
)
|
||||
|
||||
output = jsonllm()
|
||||
print(output)
|
||||
@@ -1,27 +0,0 @@
|
||||
from jsonschema import validate
|
||||
|
||||
simple_weather_schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"humidity": {"type": "number"},
|
||||
"temperatureC": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"value": {"type": "number"},
|
||||
"unit": {"type": "string", "enum": ["celsius", "fahrenheit"]},
|
||||
},
|
||||
"required": ["value", "unit"],
|
||||
},
|
||||
},
|
||||
"required": ["humidity", "temperature"],
|
||||
}
|
||||
|
||||
validate(
|
||||
{
|
||||
"humidity": 0.9,
|
||||
"temperature": {
|
||||
"value": 37,
|
||||
},
|
||||
},
|
||||
simple_weather_schema,
|
||||
)
|
||||
-291
@@ -1,291 +0,0 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 332,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"from transformers import (\n",
|
||||
" AutoModelForCausalLM,\n",
|
||||
" AutoTokenizer,\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"model_name = \"EleutherAI/gpt-neo-1.3B\"\n",
|
||||
"model = AutoModelForCausalLM.from_pretrained(\n",
|
||||
" model_name,\n",
|
||||
")\n",
|
||||
"tokenizer = AutoTokenizer.from_pretrained(model_name)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 333,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"tensor([[2579, 292, 7568]])\n",
|
||||
" 28\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"to_tokenize = \" 28asdf\"\n",
|
||||
"input_ids = tokenizer.encode(to_tokenize, return_tensors=\"pt\")\n",
|
||||
"print(input_ids)\n",
|
||||
"\n",
|
||||
"to_decode = [2579]\n",
|
||||
"print(tokenizer.decode(to_decode, skip_special_tokens=True))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 334,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from transformers import (\n",
|
||||
" PreTrainedTokenizer,\n",
|
||||
" LogitsWarper,StoppingCriteria\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# class NumberStoppingCriteria(StoppingCriteria):\n",
|
||||
"# \"\"\"\n",
|
||||
"# This class can be used to stop generation when there is a repeated decimal point in the generated text.\n",
|
||||
"# \"\"\"\n",
|
||||
"\n",
|
||||
"# def __init__(self, tokenizer: PreTrainedTokenizer):\n",
|
||||
"# self.tokenizer = tokenizer\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"# def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool:\n",
|
||||
"# decoded = self.tokenizer.decode(input_ids[0], skip_special_tokens=True)\n",
|
||||
"# if \"..\" in decoded:\n",
|
||||
"# return True\n",
|
||||
" \n",
|
||||
"# if ' ' in decoded.strip():\n",
|
||||
"# return True\n",
|
||||
" \n",
|
||||
"# return False\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"class OutputNumbersTokens(LogitsWarper):\n",
|
||||
" def __init__(self, tokenizer: PreTrainedTokenizer, prompt: str):\n",
|
||||
" self.whitelist_tokens = [tokenizer.eos_token_id]\n",
|
||||
" self.tokenized_prompt = tokenizer(prompt, return_tensors=\"pt\")\n",
|
||||
"\n",
|
||||
" for token_str, token_id in tokenizer.get_vocab().items():\n",
|
||||
"\n",
|
||||
" if (\n",
|
||||
" token_str.startswith(\"Ġ\") and (\n",
|
||||
" all(c.isdigit() or c == \".\" for c in token_str[1:]) and token_str.count(\".\") <= 1\n",
|
||||
" ) or (\n",
|
||||
" all(c.isdigit() or c == \".\" for c in token_str) and token_str.count(\".\") <= 1\n",
|
||||
" )\n",
|
||||
" ):\n",
|
||||
" self.whitelist_tokens.append(token_id)\n",
|
||||
"\n",
|
||||
" def __call__(self, input_ids, scores):\n",
|
||||
" input_ids = input_ids[:, len(self.tokenized_prompt[\"input_ids\"][0]):]\n",
|
||||
" scores[\n",
|
||||
" :, [i for i in range(len(scores[0])) if i not in self.whitelist_tokens]\n",
|
||||
" ] = -float(\"inf\")\n",
|
||||
" return scores\n",
|
||||
"\n",
|
||||
"a = OutputNumbersTokens(tokenizer, prompt)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 326,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"found space28 2579\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"\n",
|
||||
"a = OutputNumbersTokens(tokenizer, prompt)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 331,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"The attention mask and the pad token id were not set. As a consequence, you may observe unexpected behavior. Please pass your input's `attention_mask` to obtain reliable results.\n",
|
||||
"Setting `pad_token_id` to `eos_token_id`:50256 for open-end generation.\n",
|
||||
"The attention mask and the pad token id were not set. As a consequence, you may observe unexpected behavior. Please pass your input's `attention_mask` to obtain reliable results.\n",
|
||||
"Setting `pad_token_id` to `eos_token_id`:50256 for open-end generation.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
" 65 tensor(-3.7568)\n",
|
||||
" 60 tensor(-3.8070)\n",
|
||||
" 75 tensor(-3.9135)\n",
|
||||
" 80 tensor(-4.0881)\n",
|
||||
" 87 tensor(-4.1663)\n",
|
||||
" 82 tensor(-4.2518)\n",
|
||||
" 50 tensor(-4.2757)\n",
|
||||
" 5 tensor(-4.3758)\n",
|
||||
" 64 tensor(-4.3760)\n",
|
||||
" 72 tensor(-4.4309)\n",
|
||||
"tensor([[ 1820, 49890, 338, 2479, 287, 812, 318, 6135]])\n",
|
||||
"my grandma's age in years is 65\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"The attention mask and the pad token id were not set. As a consequence, you may observe unexpected behavior. Please pass your input's `attention_mask` to obtain reliable results.\n",
|
||||
"Setting `pad_token_id` to `eos_token_id`:50256 for open-end generation.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
" 0 tensor(-3.4328)\n",
|
||||
" 5 tensor(-4.0209)\n",
|
||||
" 36 tensor(-4.2786)\n",
|
||||
" 19 tensor(-4.4433)\n",
|
||||
" 59 tensor(-4.5570)\n",
|
||||
" 24 tensor(-4.7431)\n",
|
||||
" 33 tensor(-4.7621)\n",
|
||||
" 56 tensor(-4.7745)\n",
|
||||
" 47 tensor(-4.8474)\n",
|
||||
" 53 tensor(-4.8719)\n",
|
||||
"tensor([[40838, 338, 5951, 287, 7370, 269, 1424, 318, 657]])\n",
|
||||
"today's temperature in degrees cels is 0\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stderr",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"The attention mask and the pad token id were not set. As a consequence, you may observe unexpected behavior. Please pass your input's `attention_mask` to obtain reliable results.\n",
|
||||
"Setting `pad_token_id` to `eos_token_id`:50256 for open-end generation.\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
" 0 tensor(-5.0538)\n",
|
||||
" 5 tensor(-6.4346)\n",
|
||||
" 7 tensor(-7.2013)\n",
|
||||
" 60 tensor(-7.4283)\n",
|
||||
" 24 tensor(-7.6246)\n",
|
||||
" 42 tensor(-7.6409)\n",
|
||||
" 50 tensor(-7.6574)\n",
|
||||
" 30 tensor(-7.6641)\n",
|
||||
" 18 tensor(-7.7838)\n",
|
||||
" 59 tensor(-7.8880)\n",
|
||||
"tensor([[40838, 338, 5951, 287, 277, 318, 657]])\n",
|
||||
"today's temperature in f is 0\n",
|
||||
" 0 tensor(-1.6487)\n",
|
||||
" 5 tensor(-3.0234)\n",
|
||||
" 7 tensor(-4.6462)\n",
|
||||
" 16 tensor(-5.9200)\n",
|
||||
" 18 tensor(-6.3192)\n",
|
||||
" 19 tensor(-6.6764)\n",
|
||||
" 24 tensor(-6.7219)\n",
|
||||
" 30 tensor(-7.0326)\n",
|
||||
"2 tensor(-7.3029)\n",
|
||||
" 50 tensor(-7.7194)\n",
|
||||
"tensor([[ 16, 1343, 352, 796, 657]])\n",
|
||||
"1 + 1 = 0\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"prompts = [\"my grandma's age in years is\", \"today's temperature in degrees cels is\", \"today's temperature in f is\", \n",
|
||||
" \"1 + 1 =\"]\n",
|
||||
"\n",
|
||||
"for prompt in prompts:\n",
|
||||
" number = model.generate(\n",
|
||||
" tokenizer.encode(prompt, return_tensors=\"pt\"),\n",
|
||||
" max_new_tokens=5,\n",
|
||||
" num_return_sequences=1,\n",
|
||||
" logits_processor=[a],\n",
|
||||
" stopping_criteria=[NumberStoppingCriteria(tokenizer)],\n",
|
||||
" ) \n",
|
||||
"\n",
|
||||
" print(number)\n",
|
||||
"\n",
|
||||
" decoded_output = tokenizer.decode(number[0], skip_special_tokens=True)\n",
|
||||
"\n",
|
||||
" print(decoded_output.strip())\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "jsonllm-4DT72Dd8-py3.10",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.10.6"
|
||||
},
|
||||
"orig_nbformat": 4
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 2
|
||||
}
|
||||
Generated
+30
-2
@@ -1,4 +1,4 @@
|
||||
# This file is automatically @generated by Poetry and should not be changed by hand.
|
||||
# This file is automatically @generated by Poetry 1.4.2 and should not be changed by hand.
|
||||
|
||||
[[package]]
|
||||
name = "accelerate"
|
||||
@@ -1545,6 +1545,21 @@ files = [
|
||||
[package.dependencies]
|
||||
mpmath = ">=0.19"
|
||||
|
||||
[[package]]
|
||||
name = "termcolor"
|
||||
version = "2.3.0"
|
||||
description = "ANSI color formatting for output in terminal"
|
||||
category = "main"
|
||||
optional = false
|
||||
python-versions = ">=3.7"
|
||||
files = [
|
||||
{file = "termcolor-2.3.0-py3-none-any.whl", hash = "sha256:3afb05607b89aed0ffe25202399ee0867ad4d3cb4180d98aaf8eefa6a5f7d475"},
|
||||
{file = "termcolor-2.3.0.tar.gz", hash = "sha256:b5b08f68937f138fe92f6c089b99f1e2da0ae56c52b78bf7075fd95420fd9a5a"},
|
||||
]
|
||||
|
||||
[package.extras]
|
||||
tests = ["pytest", "pytest-cov"]
|
||||
|
||||
[[package]]
|
||||
name = "tokenizers"
|
||||
version = "0.13.3"
|
||||
@@ -1608,6 +1623,10 @@ category = "main"
|
||||
optional = false
|
||||
python-versions = ">=3.8.0"
|
||||
files = [
|
||||
{file = "torch-2.0.0-1-cp310-cp310-manylinux2014_aarch64.whl", hash = "sha256:c9090bda7d2eeeecd74f51b721420dbeb44f838d4536cc1b284e879417e3064a"},
|
||||
{file = "torch-2.0.0-1-cp311-cp311-manylinux2014_aarch64.whl", hash = "sha256:bd42db2a48a20574d2c33489e120e9f32789c4dc13c514b0c44272972d14a2d7"},
|
||||
{file = "torch-2.0.0-1-cp38-cp38-manylinux2014_aarch64.whl", hash = "sha256:8969aa8375bcbc0c2993e7ede0a7f889df9515f18b9b548433f412affed478d9"},
|
||||
{file = "torch-2.0.0-1-cp39-cp39-manylinux2014_aarch64.whl", hash = "sha256:ab2da16567cb55b67ae39e32d520d68ec736191d88ac79526ca5874754c32203"},
|
||||
{file = "torch-2.0.0-cp310-cp310-manylinux1_x86_64.whl", hash = "sha256:7a9319a67294ef02459a19738bbfa8727bb5307b822dadd708bc2ccf6c901aca"},
|
||||
{file = "torch-2.0.0-cp310-cp310-manylinux2014_aarch64.whl", hash = "sha256:9f01fe1f6263f31bd04e1757946fd63ad531ae37f28bb2dbf66f5c826ee089f4"},
|
||||
{file = "torch-2.0.0-cp310-cp310-win_amd64.whl", hash = "sha256:527f4ae68df7b8301ee6b1158ca56350282ea633686537b30dbb5d7b4a52622a"},
|
||||
@@ -1786,6 +1805,15 @@ category = "main"
|
||||
optional = false
|
||||
python-versions = "*"
|
||||
files = [
|
||||
{file = "triton-2.0.0-1-cp310-cp310-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:38806ee9663f4b0f7cd64790e96c579374089e58f49aac4a6608121aa55e2505"},
|
||||
{file = "triton-2.0.0-1-cp311-cp311-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:226941c7b8595219ddef59a1fdb821e8c744289a132415ddd584facedeb475b1"},
|
||||
{file = "triton-2.0.0-1-cp36-cp36m-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:4c9fc8c89874bc48eb7e7b2107a9b8d2c0bf139778637be5bfccb09191685cfd"},
|
||||
{file = "triton-2.0.0-1-cp37-cp37m-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:d2684b6a60b9f174f447f36f933e9a45f31db96cb723723ecd2dcfd1c57b778b"},
|
||||
{file = "triton-2.0.0-1-cp38-cp38-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:9d4978298b74fcf59a75fe71e535c092b023088933b2f1df933ec32615e4beef"},
|
||||
{file = "triton-2.0.0-1-cp39-cp39-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:74f118c12b437fb2ca25e1a04759173b517582fcf4c7be11913316c764213656"},
|
||||
{file = "triton-2.0.0-1-pp37-pypy37_pp73-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:9618815a8da1d9157514f08f855d9e9ff92e329cd81c0305003eb9ec25cc5add"},
|
||||
{file = "triton-2.0.0-1-pp38-pypy38_pp73-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:1aca3303629cd3136375b82cb9921727f804e47ebee27b2677fef23005c3851a"},
|
||||
{file = "triton-2.0.0-1-pp39-pypy39_pp73-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:e3e13aa8b527c9b642e3a9defcc0fbd8ffbe1c80d8ac8c15a01692478dc64d8a"},
|
||||
{file = "triton-2.0.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:8f05a7e64e4ca0565535e3d5d3405d7e49f9d308505bb7773d21fb26a4c008c2"},
|
||||
{file = "triton-2.0.0-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:bb4b99ca3c6844066e516658541d876c28a5f6e3a852286bbc97ad57134827fd"},
|
||||
{file = "triton-2.0.0-cp36-cp36m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:47b4d70dc92fb40af553b4460492c31dc7d3a114a979ffb7a5cdedb7eb546c08"},
|
||||
@@ -1866,4 +1894,4 @@ test = ["pytest (>=6.0.0)"]
|
||||
[metadata]
|
||||
lock-version = "2.0"
|
||||
python-versions = "^3.10"
|
||||
content-hash = "0b4d90066aab259321542f8407d3cb8d907bcdbcd01301cf7b3c58e50f8d45d6"
|
||||
content-hash = "b57850a2d5b4340321c5c7f0c5cfffd843b22d84e71c6c560743a761d91da612"
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
[virtualenvs]
|
||||
in-project = true
|
||||
+7
-2
@@ -1,8 +1,8 @@
|
||||
[tool.poetry]
|
||||
name = "jsonllm"
|
||||
name = "jsonformer"
|
||||
version = "0.1.0"
|
||||
description = ""
|
||||
authors = ["rahulgs12 <gsr1998@gmail.com>"]
|
||||
authors = ["1rgs <rgsduke@gmail.com>"]
|
||||
readme = "README.md"
|
||||
|
||||
[tool.poetry.dependencies]
|
||||
@@ -12,6 +12,7 @@ jsonschema = "^4.17.3"
|
||||
torch = "^2.0.0"
|
||||
accelerate = "^0.18.0"
|
||||
bitsandbytes = "^0.38.1"
|
||||
termcolor = "^2.3.0"
|
||||
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
@@ -20,3 +21,7 @@ ipykernel = "^6.22.0"
|
||||
[build-system]
|
||||
requires = ["poetry-core"]
|
||||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[virtualenvs]
|
||||
create = true
|
||||
in-project = true
|
||||
@@ -1 +1,74 @@
|
||||
# jsonllm
|
||||
# Jsonformer: A Bulletproof Way to Generate Structured JSON from Language Models.
|
||||
|
||||
## Problem: Getting models to output structed JSON is hard
|
||||
|
||||
## Solution: Only generate the content tokens and fill in the fixed tokens
|
||||
|
||||

|
||||
|
||||
Generating structured JSON from language models is a challenging task. The
|
||||
generated JSON must be syntactically correct, and it must conform to a schema
|
||||
that specifies the structure of the JSON.
|
||||
|
||||
Current approaches to this problem are brittle and error-prone. They rely on prompt engineering, fine-tuning, and post-processing, but they still fail to generate syntactically correct JSON in many cases.
|
||||
|
||||
## Example
|
||||
|
||||
```python
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
model = AutoModelForCausalLM.from_pretrained("databricks/dolly-v2-12b")
|
||||
tokenizer = AutoTokenizer.from_pretrained("databricks/dolly-v2-12b")
|
||||
|
||||
schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"age": {"type": "number"},
|
||||
"is_student": {"type": "boolean"},
|
||||
"courses": {
|
||||
"type": "array",
|
||||
"items": {"type": "string"}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
prompt = "Generate a person's information based on the following schema:"
|
||||
jsonformer = Jsonformer(model, tokenizer, json_schema, prompt)
|
||||
generated_data = jsonformer()
|
||||
|
||||
print(generated_data)
|
||||
```
|
||||
|
||||
Jsonformer is a new approach to this problem. In structured data, many tokens are fixed and predictable. Jsonformer is a wrapper around HuggingFace models that fills in the fixed tokens during the generation process, and only delegates the generation of content tokens to the language model. This makes it more efficient and bulletproof than existing approaches.
|
||||
|
||||
This currently supports a subset of JSON Schema. Below is a list of the supported schema types:
|
||||
|
||||
- number
|
||||
- boolean
|
||||
- string
|
||||
- array
|
||||
- object
|
||||
|
||||
## Features
|
||||
|
||||
- Bulletproof JSON generation: Jsonformer ensures that the generated JSON is always syntactically correct and conforms to the specified schema.
|
||||
- Efficiency: By generating only the content tokens and filling in the fixed tokens, Jsonformer is more efficient than generating a full JSON string and parsing it.
|
||||
- Flexible and extendable: Jsonformer is built on top of the HuggingFace transformers library, making it compatible with any model that supports the HuggingFace interface.
|
||||
|
||||
## Usage
|
||||
|
||||
To use Jsonformer, you need to provide a language model, a tokenizer, a JSON schema, and a prompt. Optionally, you can enable debugging, and set the maximum array length for generated arrays.
|
||||
|
||||
```python
|
||||
from transformers import PreTrainedModel, PreTrainedTokenizer
|
||||
from typing import Dict, Any
|
||||
from jsonformer import Jsonformer
|
||||
|
||||
jsonformer = Jsonformer(model, tokenizer, json_schema, prompt, debug=True, max_array_length=10)
|
||||
generated_data = jsonformer()
|
||||
```
|
||||
|
||||
## License
|
||||
|
||||
Jsonformer is released under the MIT License. You are free to use, modify, and distribute this software for any purpose, commercial or non-commercial, as long as the original copyright and license notice are included.
|
||||
|
||||
@@ -1,258 +0,0 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 2,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import json\n",
|
||||
"import random\n",
|
||||
"from transformers import (\n",
|
||||
" PreTrainedTokenizer,\n",
|
||||
" PreTrainedModel,\n",
|
||||
" AutoModelForCausalLM,\n",
|
||||
" AutoTokenizer,\n",
|
||||
")\n",
|
||||
"from typing import Any, Dict\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"model_name = \"databricks/dolly-v2-3b\"\n",
|
||||
"model = AutoModelForCausalLM.from_pretrained(\n",
|
||||
" model_name,\n",
|
||||
")\n",
|
||||
"tokenizer = AutoTokenizer.from_pretrained(model_name)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"data": {
|
||||
"text/plain": [
|
||||
"<jsonllm.logits_processors.NumberStoppingCriteria at 0x1777c4d60>"
|
||||
]
|
||||
},
|
||||
"execution_count": 3,
|
||||
"metadata": {},
|
||||
"output_type": "execute_result"
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"from jsonllm.logits_processors import NumberStoppingCriteria, OutputNumbersTokens\n",
|
||||
"\n",
|
||||
"NumberStoppingCriteria(tokenizer, 2)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 10,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"class JSONLLM:\n",
|
||||
" value: Dict[str, Any] = {}\n",
|
||||
"\n",
|
||||
" def __init__(\n",
|
||||
" self,\n",
|
||||
" model: PreTrainedModel,\n",
|
||||
" tokenizer: PreTrainedTokenizer,\n",
|
||||
" json_schema: Dict[str, Any],\n",
|
||||
" prompt: str,\n",
|
||||
" ):\n",
|
||||
" self.model = model\n",
|
||||
" self.tokenizer = tokenizer\n",
|
||||
" self.json_schema = json_schema\n",
|
||||
" self.prompt = prompt\n",
|
||||
"\n",
|
||||
" self.number_logit_processor = OutputNumbersTokens(self.tokenizer, self.prompt)\n",
|
||||
" self.number_stop_criteria = NumberStoppingCriteria(self.tokenizer, 3)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" def generate_number(self, suffix=\"\") -> float:\n",
|
||||
" print(\"generate_number\", suffix)\n",
|
||||
" prompt = self.get_prompt() + suffix\n",
|
||||
"\n",
|
||||
" print(\"\\033[91m {}\\033[00m\".format(prompt))\n",
|
||||
" response = self.model.generate(\n",
|
||||
" self.tokenizer.encode(prompt, return_tensors=\"pt\"),\n",
|
||||
" max_new_tokens=6,\n",
|
||||
" num_return_sequences=1,\n",
|
||||
" logits_processor=[self.number_logit_processor],\n",
|
||||
" stopping_criteria=[self.number_stop_criteria],\n",
|
||||
" temperature=1.5,\n",
|
||||
" pad_token_id=tokenizer.eos_token_id\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" response = self.tokenizer.decode(response[0], skip_special_tokens=True)\n",
|
||||
" print(\"response is\")\n",
|
||||
" print(\"\\033[94m {}\\033[00m\".format(response))\n",
|
||||
" response = response.strip().rstrip(\".\").lstrip(\"0\")\n",
|
||||
" try:\n",
|
||||
" return float(response)\n",
|
||||
" except ValueError:\n",
|
||||
" print(\"ValueError\")\n",
|
||||
" return \n",
|
||||
"\n",
|
||||
" def generate_boolean(self, suffix=\"\") -> bool:\n",
|
||||
" prompt = self.get_prompt()\n",
|
||||
" true_token_id = self.tokenizer.encode(\"true\", add_special_tokens=False)[0]\n",
|
||||
" false_token_id = self.tokenizer.encode(\"false\", add_special_tokens=False)[0]\n",
|
||||
"\n",
|
||||
" response = self.generate(\n",
|
||||
" prompt, forced_bos_token_id=[true_token_id, false_token_id]\n",
|
||||
" ).lower()\n",
|
||||
"\n",
|
||||
" if response == \"true\":\n",
|
||||
" return True\n",
|
||||
" else:\n",
|
||||
" return False\n",
|
||||
"\n",
|
||||
" def generate_array(\n",
|
||||
" self, item_schema: Dict[str, Any], obj: Dict[str, Any], suffix=\"\"\n",
|
||||
" ) -> list:\n",
|
||||
" array_length = random.randint(0, 5)\n",
|
||||
" return [self.generate_value(item_schema, obj) for _ in range(array_length)]\n",
|
||||
"\n",
|
||||
" # add stopping criteria with \"\n",
|
||||
" def generate_string(self) -> str:\n",
|
||||
" prompt = self.get_prompt()\n",
|
||||
" response = self.generate(prompt)\n",
|
||||
" return response\n",
|
||||
"\n",
|
||||
" def generate_object(\n",
|
||||
" self, properties: Dict[str, Any], obj: Dict[str, Any], suffix=\"\"\n",
|
||||
" ) -> Dict[str, Any]:\n",
|
||||
" print(\"generate_object\", properties)\n",
|
||||
"\n",
|
||||
" for key, schema in properties.items():\n",
|
||||
" value = self.generate_value(schema, obj, suffix=f'\"{key}\": ')\n",
|
||||
"\n",
|
||||
" obj[key] = value\n",
|
||||
" return obj\n",
|
||||
"\n",
|
||||
" def generate_value(self, schema: Dict[str, Any], obj: Dict[str, Any], suffix=\"\"):\n",
|
||||
" schema_type = schema[\"type\"]\n",
|
||||
" if schema_type == \"number\":\n",
|
||||
" return self.generate_number(suffix=suffix)\n",
|
||||
" elif schema_type == \"boolean\":\n",
|
||||
" return self.generate_boolean(suffix=suffix)\n",
|
||||
" elif schema_type == \"array\":\n",
|
||||
" return self.generate_array(schema[\"items\"], obj, suffix=suffix)\n",
|
||||
" elif schema_type == \"object\":\n",
|
||||
" return self.generate_object(schema[\"properties\"], obj, suffix=suffix)\n",
|
||||
" else:\n",
|
||||
" raise ValueError(f\"Unsupported schema type: {schema_type}\")\n",
|
||||
"\n",
|
||||
" def get_prompt(self):\n",
|
||||
" template = \"\"\"{prompt}\\nMake sure to output in the following format:\\n{schema}\\n Result: {progress}\"\"\"\n",
|
||||
" progress = json.dumps(self.value)\n",
|
||||
"\n",
|
||||
" progress = progress.rstrip(\"}\").rstrip(\"]\").rstrip(\",\")\n",
|
||||
"\n",
|
||||
" prompt = template.format(\n",
|
||||
" prompt=self.prompt,\n",
|
||||
" schema=json.dumps(self.json_schema),\n",
|
||||
" progress=progress,\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" return prompt\n",
|
||||
"\n",
|
||||
" def __call__(self) -> Dict[str, Any]:\n",
|
||||
" self.value = {}\n",
|
||||
" generated_data = self.generate_object(\n",
|
||||
" self.json_schema[\"properties\"], self.value\n",
|
||||
" )\n",
|
||||
" return generated_data\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 11,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"generate_object {'temperature': {'type': 'number'}}\n",
|
||||
"generate_number \"temperature\": \n",
|
||||
"\u001b[91m Generate a weather object\n",
|
||||
"Make sure to output in the following format:\n",
|
||||
"{\"type\": \"object\", \"properties\": {\"temperature\": {\"type\": \"number\"}}}\n",
|
||||
" Result: {\"temperature\": \u001b[00m\n",
|
||||
"Stopping because of multiple .\n",
|
||||
"response is\n",
|
||||
"\u001b[94m Generate a weather object\n",
|
||||
"Make sure to output in the following format:\n",
|
||||
"{\"type\": \"object\", \"properties\": {\"temperature\": {\"type\": \"number\"}}}\n",
|
||||
" Result: {\"temperature\": 000000000.0.\u001b[00m\n",
|
||||
"ValueError\n",
|
||||
"{'temperature': None}\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"source": [
|
||||
"weather_schema = {\n",
|
||||
" \"type\": \"object\",\n",
|
||||
" \"properties\": {\n",
|
||||
" \"temperature\": {\"type\": \"number\"},\n",
|
||||
" # \"humidity\": {\"type\": \"number\"},\n",
|
||||
" },\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"jsonllm = JSONLLM(\n",
|
||||
" model=model,\n",
|
||||
" tokenizer=tokenizer,\n",
|
||||
" json_schema=weather_schema,\n",
|
||||
" prompt=\"Generate a weather object\",\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"output = jsonllm()\n",
|
||||
"print(output)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "jsonllm-4DT72Dd8-py3.10",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.10.6"
|
||||
},
|
||||
"orig_nbformat": 4
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 2
|
||||
}
|
||||
Reference in New Issue
Block a user