From a891783f53c0ae094464fbc981db772208d3d609 Mon Sep 17 00:00:00 2001 From: Ryan <18477649+Ryul0rd@users.noreply.github.com> Date: Sun, 14 May 2023 21:51:36 -0700 Subject: [PATCH] add integer type --- jsonformer/logits_processors.py | 51 +++++++++++++++++++++++++++++++++ jsonformer/main.py | 39 +++++++++++++++++++++++++ 2 files changed, 90 insertions(+) diff --git a/jsonformer/logits_processors.py b/jsonformer/logits_processors.py index db288d3..c1088ce 100644 --- a/jsonformer/logits_processors.py +++ b/jsonformer/logits_processors.py @@ -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 diff --git a/jsonformer/main.py b/jsonformer/main.py index dd867d4..8b5ac8e 100644 --- a/jsonformer/main.py +++ b/jsonformer/main.py @@ -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