mirror of
https://github.com/wassname/prob_jsonformer.git
synced 2026-09-09 11:29:57 +08:00
improve number decoding
This commit is contained in:
@@ -47,7 +47,7 @@ json_schema = {
|
||||
}
|
||||
|
||||
prompt = "Generate a person's information based on the following schema:"
|
||||
jsonformer = Jsonformer(model, tokenizer, json_schema, prompt, device="cuda")
|
||||
jsonformer = Jsonformer(model, tokenizer, json_schema, prompt)
|
||||
generated_data = jsonformer()
|
||||
|
||||
print(generated_data)
|
||||
|
||||
+100582
-23
File diff suppressed because it is too large
Load Diff
@@ -38,8 +38,6 @@ builder = Jsonformer(
|
||||
tokenizer=tokenizer,
|
||||
json_schema=weather_schema,
|
||||
prompt="generate the weather",
|
||||
debug=True,
|
||||
device="cuda",
|
||||
)
|
||||
|
||||
print("Generating...")
|
||||
|
||||
@@ -1,18 +1,28 @@
|
||||
from typing import List
|
||||
from transformers import PreTrainedTokenizer, LogitsWarper, StoppingCriteria
|
||||
import torch
|
||||
|
||||
|
||||
class NumberStoppingCriteria(StoppingCriteria):
|
||||
def __init__(self, tokenizer: PreTrainedTokenizer, precision: int = 2):
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: PreTrainedTokenizer,
|
||||
prompt_length: int,
|
||||
precision: int = 3,
|
||||
):
|
||||
self.tokenizer = tokenizer
|
||||
self.precision = precision
|
||||
self.prompt_length = prompt_length
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
input_ids: torch.LongTensor,
|
||||
scores: torch.FloatTensor,
|
||||
) -> bool:
|
||||
decoded = self.tokenizer.decode(input_ids[0], skip_special_tokens=True)
|
||||
decoded = self.tokenizer.decode(
|
||||
input_ids[0][self.prompt_length :], skip_special_tokens=True
|
||||
)
|
||||
|
||||
if decoded.count(".") > 1:
|
||||
return True
|
||||
|
||||
@@ -22,38 +32,35 @@ class NumberStoppingCriteria(StoppingCriteria):
|
||||
):
|
||||
return True
|
||||
|
||||
if (
|
||||
len(decoded) > 1
|
||||
and any(c.isdigit() for c in decoded)
|
||||
and decoded[-1] in [" ", "\n"]
|
||||
):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
class OutputNumbersTokens(LogitsWarper):
|
||||
def __init__(self, tokenizer: PreTrainedTokenizer, prompt: str):
|
||||
self.whitelist_tokens = [
|
||||
# tokenizer.eos_token_id
|
||||
]
|
||||
self.whitelist_tokens = []
|
||||
self.tokenized_prompt = tokenizer(prompt, return_tensors="pt")
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
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])
|
||||
)
|
||||
for _, token_id in tokenizer.get_vocab().items():
|
||||
token_str = tokenizer.decode(token_id)
|
||||
token_str = token_str.strip()
|
||||
|
||||
if token_str == "" or (
|
||||
all(c.isdigit() or c == "." for c in token_str)
|
||||
and token_str.count(".") <= 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")
|
||||
|
||||
+14
-9
@@ -17,7 +17,6 @@ class Jsonformer:
|
||||
json_schema: Dict[str, Any],
|
||||
prompt: str,
|
||||
*,
|
||||
device: str,
|
||||
debug: bool = False,
|
||||
max_array_length: int = 10,
|
||||
max_number_tokens: int = 6,
|
||||
@@ -30,7 +29,6 @@ class Jsonformer:
|
||||
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
|
||||
@@ -39,22 +37,26 @@ 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:
|
||||
print(*args, **kwargs)
|
||||
|
||||
def generate_number(self) -> float:
|
||||
def generate_number(self, temperature: Union[float, None] = None, iterations=0):
|
||||
prompt = self.get_prompt()
|
||||
self.debug("[generate_number] prompt", prompt)
|
||||
input_tokens = self.tokenizer.encode(prompt, return_tensors="pt").to(
|
||||
self.model.device
|
||||
)
|
||||
response = self.model.generate(
|
||||
self.tokenizer.encode(prompt, return_tensors="pt").to(self.model.device),
|
||||
input_tokens,
|
||||
max_new_tokens=self.max_number_tokens,
|
||||
num_return_sequences=1,
|
||||
logits_processor=[self.number_logit_processor],
|
||||
stopping_criteria=[self.number_stop_criteria],
|
||||
temperature=self.temperature,
|
||||
stopping_criteria=[
|
||||
NumberStoppingCriteria(self.tokenizer, len(input_tokens[0]))
|
||||
],
|
||||
temperature=temperature or self.temperature,
|
||||
pad_token_id=self.tokenizer.eos_token_id,
|
||||
)
|
||||
response = self.tokenizer.decode(response[0], skip_special_tokens=True)
|
||||
@@ -62,11 +64,14 @@ class Jsonformer:
|
||||
response = response[len(prompt) :]
|
||||
response = response.strip().rstrip(".")
|
||||
|
||||
print("response", "|" + response + "|")
|
||||
try:
|
||||
return float(response)
|
||||
except ValueError:
|
||||
print("ValueError")
|
||||
return
|
||||
if iterations > 3:
|
||||
raise ValueError("Failed to generate a valid number")
|
||||
|
||||
return self.generate_number(temperature=self.temperature * 1.3)
|
||||
|
||||
def generate_boolean(self) -> bool:
|
||||
prompt = self.get_prompt()
|
||||
|
||||
Reference in New Issue
Block a user