diff --git a/.gitignore b/.gitignore index 3104065..197d512 100644 --- a/.gitignore +++ b/.gitignore @@ -1 +1,3 @@ -jsonllm/__pycache__ \ No newline at end of file +jsonllm/__pycache__ +.venv +workspace.ipynb diff --git a/example.ipynb b/example.ipynb new file mode 100644 index 0000000..8a5919d --- /dev/null +++ b/example.ipynb @@ -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 +} diff --git a/img/cover.png b/img/cover.png new file mode 100644 index 0000000..415250c Binary files /dev/null and b/img/cover.png differ diff --git a/img/cover2.png b/img/cover2.png new file mode 100644 index 0000000..7420674 Binary files /dev/null and b/img/cover2.png differ diff --git a/img/cover3.png b/img/cover3.png new file mode 100644 index 0000000..ec303e5 Binary files /dev/null and b/img/cover3.png differ diff --git a/jsonformer/__pycache__/example.cpython-310.pyc b/jsonformer/__pycache__/example.cpython-310.pyc new file mode 100644 index 0000000..6467cd1 Binary files /dev/null and b/jsonformer/__pycache__/example.cpython-310.pyc differ diff --git a/jsonformer/__pycache__/format.cpython-310.pyc b/jsonformer/__pycache__/format.cpython-310.pyc new file mode 100644 index 0000000..df49d65 Binary files /dev/null and b/jsonformer/__pycache__/format.cpython-310.pyc differ diff --git a/jsonformer/__pycache__/logits_processors.cpython-310.pyc b/jsonformer/__pycache__/logits_processors.cpython-310.pyc new file mode 100644 index 0000000..be87112 Binary files /dev/null and b/jsonformer/__pycache__/logits_processors.cpython-310.pyc differ diff --git a/jsonformer/__pycache__/main.cpython-310.pyc b/jsonformer/__pycache__/main.cpython-310.pyc new file mode 100644 index 0000000..f935095 Binary files /dev/null and b/jsonformer/__pycache__/main.cpython-310.pyc differ diff --git a/jsonformer/example.py b/jsonformer/example.py new file mode 100644 index 0000000..72f0c8f --- /dev/null +++ b/jsonformer/example.py @@ -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, +) diff --git a/jsonformer/format.py b/jsonformer/format.py new file mode 100644 index 0000000..8d19943 --- /dev/null +++ b/jsonformer/format.py @@ -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) diff --git a/jsonllm/logits_processors.py b/jsonformer/logits_processors.py similarity index 81% rename from jsonllm/logits_processors.py rename to jsonformer/logits_processors.py index 53bb8da..eb0e884 100644 --- a/jsonllm/logits_processors.py +++ b/jsonformer/logits_processors.py @@ -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(): diff --git a/jsonformer/main.py b/jsonformer/main.py new file mode 100644 index 0000000..2bfd3f6 --- /dev/null +++ b/jsonformer/main.py @@ -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 diff --git a/jsonllm/main.py b/jsonllm/main.py deleted file mode 100644 index 8aeedcf..0000000 --- a/jsonllm/main.py +++ /dev/null @@ -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) diff --git a/jsonllm/test.py b/jsonllm/test.py deleted file mode 100644 index 29d85b6..0000000 --- a/jsonllm/test.py +++ /dev/null @@ -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, -) diff --git a/play.ipynb b/play.ipynb deleted file mode 100644 index 68ccac1..0000000 --- a/play.ipynb +++ /dev/null @@ -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 -} diff --git a/poetry.lock b/poetry.lock index 2721198..3256e92 100644 --- a/poetry.lock +++ b/poetry.lock @@ -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" diff --git a/poetry.toml b/poetry.toml new file mode 100644 index 0000000..ab1033b --- /dev/null +++ b/poetry.toml @@ -0,0 +1,2 @@ +[virtualenvs] +in-project = true diff --git a/pyproject.toml b/pyproject.toml index 191a53f..c680ed3 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,8 +1,8 @@ [tool.poetry] -name = "jsonllm" +name = "jsonformer" version = "0.1.0" description = "" -authors = ["rahulgs12 "] +authors = ["1rgs "] 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 \ No newline at end of file diff --git a/readme.md b/readme.md index ae314aa..5b2078d 100644 --- a/readme.md +++ b/readme.md @@ -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. diff --git a/t.ipynb b/t.ipynb deleted file mode 100644 index 75780dc..0000000 --- a/t.ipynb +++ /dev/null @@ -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": [ - "" - ] - }, - "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 -}