mirror of
https://github.com/wassname/prob_jsonformer.git
synced 2026-09-09 11:29:57 +08:00
.
This commit is contained in:
@@ -0,0 +1 @@
|
||||
jsonllm/__pycache__
|
||||
@@ -0,0 +1,67 @@
|
||||
from transformers import PreTrainedTokenizer, LogitsWarper, StoppingCriteria
|
||||
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
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
input_ids: torch.LongTensor,
|
||||
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 .")
|
||||
return True
|
||||
|
||||
if (
|
||||
decoded.strip().count(".") == 1
|
||||
and len(decoded.strip().split(".")[1]) > self.precision
|
||||
):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
class OutputNumbersTokens(LogitsWarper):
|
||||
def __init__(self, tokenizer: PreTrainedTokenizer, prompt: str):
|
||||
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():
|
||||
if (
|
||||
(
|
||||
token_str.startswith("Ġ")
|
||||
and (
|
||||
all(c.isdigit() or c == "." for c in token_str[1:])
|
||||
and token_str.count(".") <= 1
|
||||
)
|
||||
)
|
||||
or (
|
||||
all(c.isdigit() or c == "." for c in token_str)
|
||||
and token_str.count(".") <= 1
|
||||
)
|
||||
or (
|
||||
token_str[-1] == " "
|
||||
and all(c.isdigit() or c == "." for c in token_str[:-1])
|
||||
)
|
||||
):
|
||||
self.whitelist_tokens.append(token_id)
|
||||
|
||||
def __call__(self, input_ids, scores):
|
||||
input_ids = input_ids[:, len(self.tokenized_prompt["input_ids"][0]) :]
|
||||
scores[
|
||||
:, [i for i in range(len(scores[0])) if i not in self.whitelist_tokens]
|
||||
] = -float("inf")
|
||||
return scores
|
||||
+173
@@ -0,0 +1,173 @@
|
||||
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)
|
||||
@@ -0,0 +1,27 @@
|
||||
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
@@ -0,0 +1,291 @@
|
||||
{
|
||||
"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
+1869
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,22 @@
|
||||
[tool.poetry]
|
||||
name = "jsonllm"
|
||||
version = "0.1.0"
|
||||
description = ""
|
||||
authors = ["rahulgs12 <gsr1998@gmail.com>"]
|
||||
readme = "README.md"
|
||||
|
||||
[tool.poetry.dependencies]
|
||||
python = "^3.10"
|
||||
transformers = "^4.27.4"
|
||||
jsonschema = "^4.17.3"
|
||||
torch = "^2.0.0"
|
||||
accelerate = "^0.18.0"
|
||||
bitsandbytes = "^0.38.1"
|
||||
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
ipykernel = "^6.22.0"
|
||||
|
||||
[build-system]
|
||||
requires = ["poetry-core"]
|
||||
build-backend = "poetry.core.masonry.api"
|
||||
@@ -0,0 +1,258 @@
|
||||
{
|
||||
"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