support cuda #12

This commit is contained in:
rahulgs12
2023-05-04 12:25:16 -04:00
parent f65fbd9c2d
commit 3c06fe00c2
5 changed files with 1904 additions and 1897 deletions
+1814 -1813
View File
File diff suppressed because it is too large Load Diff
+75 -74
View File
@@ -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
},
+5 -2
View File
@@ -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
View File
@@ -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
)
+1 -1
View File
@@ -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)