This commit is contained in:
rahulgs12
2023-05-01 14:22:41 -04:00
parent b422f68162
commit 2dc8f882a2
21 changed files with 601 additions and 767 deletions
+3 -1
View File
@@ -1 +1,3 @@
jsonllm/__pycache__
jsonllm/__pycache__
.venv
workspace.ipynb
+217
View File
@@ -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
}
BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 60 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 42 KiB

BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 81 KiB

Binary file not shown.
Binary file not shown.
Binary file not shown.
+47
View File
@@ -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,
)
+24
View File
@@ -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():
+192
View File
@@ -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
View File
@@ -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)
-27
View File
@@ -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
View File
@@ -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
View File
@@ -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"
+2
View File
@@ -0,0 +1,2 @@
[virtualenvs]
in-project = true
+7 -2
View File
@@ -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
+74 -1
View File
@@ -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
![cover](img/cover3.png)
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.
-258
View File
@@ -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
}