mirror of
https://github.com/wassname/prob_jsonformer.git
synced 2026-09-09 11:29:57 +08:00
support cuda #12
This commit is contained in:
+1814
-1813
File diff suppressed because it is too large
Load Diff
+75
-74
@@ -9,7 +9,7 @@
|
||||
"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",
|
||||
"/home/ubuntu/jsonformer/.venv/lib/python3.8/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"
|
||||
]
|
||||
},
|
||||
@@ -26,12 +26,68 @@
|
||||
"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",
|
||||
"model_name = \"databricks/dolly-v2-3b\"\n",
|
||||
"model = AutoModelForCausalLM.from_pretrained(model_name, use_cache=True, device_map=\"auto\")\n",
|
||||
"tokenizer = AutoTokenizer.from_pretrained(model_name, use_fast=True, use_cache=True)\n",
|
||||
"print(\"Loaded model and tokenizer\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 4,
|
||||
"metadata": {},
|
||||
"outputs": [
|
||||
{
|
||||
"name": "stdout",
|
||||
"output_type": "stream",
|
||||
"text": [
|
||||
"Generating...\n",
|
||||
"{\n",
|
||||
" temperature: \u001b[32m\"22.0\"\u001b[0m,\n",
|
||||
" humidity: \u001b[32m\"60\"\u001b[0m,\n",
|
||||
" wind_speed: {\n",
|
||||
" value: \u001b[32m\"5\"\u001b[0m,\n",
|
||||
" unit: \u001b[32m\"mph\"\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\": \"string\"},\n",
|
||||
" \"humidity\": {\n",
|
||||
" \"type\": \"string\",\n",
|
||||
" },\n",
|
||||
" \"wind_speed\": {\n",
|
||||
" \"type\": \"object\",\n",
|
||||
" \"properties\": {\n",
|
||||
" \"value\": {\"type\": \"string\"},\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",
|
||||
" device=\"cuda\",\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"print(\"Generating...\")\n",
|
||||
"output = builder()\n",
|
||||
"\n",
|
||||
"highlight_values(output)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": 3,
|
||||
@@ -43,75 +99,11 @@
|
||||
"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",
|
||||
" make: \u001b[32m\"audi\"\u001b[0m,\n",
|
||||
" model: \u001b[32m\"model a4\"\u001b[0m,\n",
|
||||
" year: \u001b[32m1.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",
|
||||
" \u001b[32m\"blue\"\u001b[0m\n",
|
||||
" ]\n",
|
||||
"}\n"
|
||||
]
|
||||
@@ -123,11 +115,12 @@
|
||||
" \"properties\": {\n",
|
||||
" \"make\": {\"type\": \"string\"},\n",
|
||||
" \"model\": {\"type\": \"string\"},\n",
|
||||
" \"year\": {\"type\": \"number\"},\n",
|
||||
" \"year\": {\"type\": \"string\"},\n",
|
||||
" \"colors\": {\n",
|
||||
" \"type\": \"array\",\n",
|
||||
" \"items\": {\"type\": \"string\"},\n",
|
||||
" }\n",
|
||||
" },\n",
|
||||
" \"shouldUseTurnSignal\": {\"type\": \"boolean\"},\n",
|
||||
" },\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
@@ -136,6 +129,7 @@
|
||||
" tokenizer=tokenizer,\n",
|
||||
" json_schema=car,\n",
|
||||
" prompt=\"generate an example car\",\n",
|
||||
" device=\"cuda\",\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"print(\"Generating...\")\n",
|
||||
@@ -143,6 +137,13 @@
|
||||
"\n",
|
||||
"highlight_values(output)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
@@ -161,7 +162,7 @@
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.10.11"
|
||||
"version": "3.8.10"
|
||||
},
|
||||
"orig_nbformat": 4
|
||||
},
|
||||
|
||||
@@ -6,7 +6,6 @@ from transformers import (
|
||||
)
|
||||
|
||||
|
||||
|
||||
weather_schema = {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
@@ -26,7 +25,10 @@ weather_schema = {
|
||||
|
||||
print("Loading model and tokenizer...")
|
||||
model_name = "databricks/dolly-v2-12b"
|
||||
model = AutoModelForCausalLM.from_pretrained(model_name, use_cache=True)
|
||||
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
model_name, use_cache=True, device_map="auto"
|
||||
)
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_name, use_fast=True, use_cache=True)
|
||||
print("Loaded model and tokenizer")
|
||||
|
||||
@@ -37,6 +39,7 @@ builder = Jsonformer(
|
||||
json_schema=weather_schema,
|
||||
prompt="generate the weather",
|
||||
debug=True,
|
||||
device="cuda",
|
||||
)
|
||||
|
||||
print("Generating...")
|
||||
|
||||
+9
-7
@@ -16,6 +16,8 @@ class Jsonformer:
|
||||
tokenizer: PreTrainedTokenizer,
|
||||
json_schema: Dict[str, Any],
|
||||
prompt: str,
|
||||
*,
|
||||
device: str,
|
||||
debug: bool = False,
|
||||
max_array_length: int = 10,
|
||||
max_number_tokens: int = 6,
|
||||
@@ -37,6 +39,7 @@ class Jsonformer:
|
||||
self.max_number_tokens = max_number_tokens
|
||||
self.temperature = temperature
|
||||
self.max_string_token_length = max_string_token_length
|
||||
self.device = device
|
||||
|
||||
def debug(self, *args, **kwargs):
|
||||
if self.debug_on:
|
||||
@@ -46,7 +49,7 @@ class Jsonformer:
|
||||
prompt = self.get_prompt()
|
||||
self.debug("[generate_number] prompt", prompt)
|
||||
response = self.model.generate(
|
||||
self.tokenizer.encode(prompt, return_tensors="pt"),
|
||||
self.tokenizer.encode(prompt, return_tensors="pt").to(self.model.device),
|
||||
max_new_tokens=self.max_number_tokens,
|
||||
num_return_sequences=1,
|
||||
logits_processor=[self.number_logit_processor],
|
||||
@@ -70,7 +73,7 @@ class Jsonformer:
|
||||
self.debug("[generate_boolean] prompt", prompt)
|
||||
|
||||
input_tensor = self.tokenizer.encode(prompt, return_tensors="pt")
|
||||
output = self.model.forward(input_tensor)
|
||||
output = self.model.forward(input_tensor.to(self.model.device))
|
||||
logits = output.logits[0, -1]
|
||||
|
||||
true_token_id = self.tokenizer.convert_tokens_to_ids("true")
|
||||
@@ -91,7 +94,7 @@ class Jsonformer:
|
||||
prompt = self.get_prompt()
|
||||
self.debug("[generate_string] prompt", prompt)
|
||||
response = self.model.generate(
|
||||
self.tokenizer.encode(prompt, return_tensors="pt"),
|
||||
self.tokenizer.encode(prompt, return_tensors="pt").to(self.model.device),
|
||||
max_new_tokens=self.max_string_token_length,
|
||||
num_return_sequences=1,
|
||||
temperature=self.temperature,
|
||||
@@ -118,7 +121,7 @@ class Jsonformer:
|
||||
self,
|
||||
schema: Dict[str, Any],
|
||||
obj: Union[Dict[str, Any], List[Any]],
|
||||
key: str | None = None,
|
||||
key: Union[str, None] = None,
|
||||
) -> Any:
|
||||
schema_type = schema["type"]
|
||||
if schema_type == "number":
|
||||
@@ -162,14 +165,14 @@ class Jsonformer:
|
||||
input_prompt = self.get_prompt()
|
||||
obj.pop()
|
||||
input_tensor = self.tokenizer.encode(input_prompt, return_tensors="pt")
|
||||
output = self.model.forward(input_tensor)
|
||||
output = self.model.forward(input_tensor.to(self.model.device))
|
||||
logits = output.logits[0, -1]
|
||||
|
||||
close_bracket_token_id = self.tokenizer.convert_tokens_to_ids("]")
|
||||
comma_token_id = self.tokenizer.convert_tokens_to_ids(", ")
|
||||
close_bracket_logits = logits[close_bracket_token_id]
|
||||
comma_logits = logits[comma_token_id]
|
||||
|
||||
|
||||
if close_bracket_logits > comma_logits:
|
||||
break
|
||||
|
||||
@@ -194,7 +197,6 @@ class Jsonformer:
|
||||
|
||||
def __call__(self) -> Dict[str, Any]:
|
||||
self.value = {}
|
||||
|
||||
generated_data = self.generate_object(
|
||||
self.json_schema["properties"], self.value
|
||||
)
|
||||
|
||||
@@ -47,7 +47,7 @@ json_schema = {
|
||||
}
|
||||
|
||||
prompt = "Generate a person's information based on the following schema:"
|
||||
jsonformer = Jsonformer(model, tokenizer, json_schema, prompt)
|
||||
jsonformer = Jsonformer(model, tokenizer, json_schema, prompt, device="cuda")
|
||||
generated_data = jsonformer()
|
||||
|
||||
print(generated_data)
|
||||
|
||||
Reference in New Issue
Block a user