add integer type

This commit is contained in:
Ryan
2023-05-14 21:51:36 -07:00
parent f6366c2c35
commit a891783f53
2 changed files with 90 additions and 0 deletions
+51
View File
@@ -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
+39
View File
@@ -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