mirror of
https://github.com/wassname/prob_jsonformer.git
synced 2026-09-09 11:29:57 +08:00
add integer type
This commit is contained in:
@@ -82,3 +82,54 @@ class OutputNumbersTokens(LogitsWarper):
|
||||
scores[~mask] = -float("inf")
|
||||
|
||||
return scores
|
||||
|
||||
class IntegerStoppingCriteria(StoppingCriteria):
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: PreTrainedTokenizer,
|
||||
prompt_length: int,
|
||||
max_digits: int = 15,
|
||||
):
|
||||
self.tokenizer = tokenizer
|
||||
self.prompt_length = prompt_length
|
||||
self.max_digits = max_digits
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
input_ids: torch.LongTensor,
|
||||
scores: torch.FloatTensor,
|
||||
) -> bool:
|
||||
decoded = self.tokenizer.decode(
|
||||
input_ids[0][self.prompt_length :], skip_special_tokens=True
|
||||
)
|
||||
|
||||
if len(decoded.strip()) > self.max_digits:
|
||||
return True
|
||||
|
||||
if (
|
||||
len(decoded) > 1
|
||||
and any(c.isdigit() for c in decoded)
|
||||
and decoded[-1] in [" ", "\n"]
|
||||
):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
class OutputIntegersTokens(LogitsWarper):
|
||||
def __init__(self, tokenizer: PreTrainedTokenizer, prompt: str):
|
||||
self.tokenizer = tokenizer
|
||||
self.tokenized_prompt = tokenizer(prompt, return_tensors="pt")
|
||||
vocab_size = len(tokenizer)
|
||||
self.allowed_mask = torch.zeros(vocab_size, dtype=torch.bool)
|
||||
|
||||
for _, token_id in tokenizer.get_vocab().items():
|
||||
token_str = tokenizer.decode(token_id).strip()
|
||||
|
||||
if token_str == "" or all(c.isdigit() for c in token_str):
|
||||
self.allowed_mask[token_id] = True
|
||||
|
||||
def __call__(self, _, scores):
|
||||
mask = self.allowed_mask.expand_as(scores)
|
||||
scores[~mask] = -float("inf")
|
||||
|
||||
return scores
|
||||
|
||||
@@ -3,6 +3,8 @@ from typing import List, Union, Dict, Any
|
||||
from jsonformer.logits_processors import (
|
||||
NumberStoppingCriteria,
|
||||
OutputNumbersTokens,
|
||||
IntegerStoppingCriteria,
|
||||
OutputIntegersTokens,
|
||||
StringStoppingCriteria,
|
||||
)
|
||||
from termcolor import cprint
|
||||
@@ -34,6 +36,7 @@ class Jsonformer:
|
||||
self.prompt = prompt
|
||||
|
||||
self.number_logit_processor = OutputNumbersTokens(self.tokenizer, self.prompt)
|
||||
self.integer_logit_processor = OutputIntegersTokens(self.tokenizer, self.prompt)
|
||||
|
||||
self.generation_marker = "|GENERATION|"
|
||||
self.debug_on = debug
|
||||
@@ -82,6 +85,36 @@ class Jsonformer:
|
||||
|
||||
return self.generate_number(temperature=self.temperature * 1.3)
|
||||
|
||||
def generate_integer(self, temperature: Union[float, None] = None, iterations=0):
|
||||
prompt = self.get_prompt()
|
||||
self.debug("[generate_number]", prompt, is_prompt=True)
|
||||
input_tokens = self.tokenizer.encode(prompt, return_tensors="pt").to(
|
||||
self.model.device
|
||||
)
|
||||
response = self.model.generate(
|
||||
input_tokens,
|
||||
max_new_tokens=self.max_number_tokens,
|
||||
num_return_sequences=1,
|
||||
logits_processor=[self.integer_logit_processor],
|
||||
stopping_criteria=[
|
||||
IntegerStoppingCriteria(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)
|
||||
|
||||
response = response[len(prompt) :]
|
||||
response = response.strip()
|
||||
self.debug("[generate_integer]", response)
|
||||
try:
|
||||
return int(response)
|
||||
except ValueError:
|
||||
if iterations > 3:
|
||||
raise ValueError("Failed to generate a valid integer")
|
||||
|
||||
return self.generate_integer(temperature=self.temperature * 1.3)
|
||||
|
||||
def generate_boolean(self) -> bool:
|
||||
prompt = self.get_prompt()
|
||||
self.debug("[generate_boolean]", prompt, is_prompt=True)
|
||||
@@ -160,6 +193,12 @@ class Jsonformer:
|
||||
else:
|
||||
obj.append(self.generation_marker)
|
||||
return self.generate_number()
|
||||
elif schema_type == "integer":
|
||||
if key:
|
||||
obj[key] = self.generation_marker
|
||||
else:
|
||||
obj.append(self.generation_marker)
|
||||
return self.generate_integer()
|
||||
elif schema_type == "boolean":
|
||||
if key:
|
||||
obj[key] = self.generation_marker
|
||||
|
||||
Reference in New Issue
Block a user