From 58e8270537fcd98557a0768a85887c2d2c038485 Mon Sep 17 00:00:00 2001 From: Ryan <18477649+Ryul0rd@users.noreply.github.com> Date: Thu, 18 May 2023 02:33:46 -0700 Subject: [PATCH] Fix bug resulting in overly long numbers/ints --- jsonformer/logits_processors.py | 39 +++++++++++++++++++++++++++------ jsonformer/main.py | 8 +++++-- 2 files changed, 38 insertions(+), 9 deletions(-) diff --git a/jsonformer/logits_processors.py b/jsonformer/logits_processors.py index c1088ce..7b09ea2 100644 --- a/jsonformer/logits_processors.py +++ b/jsonformer/logits_processors.py @@ -48,14 +48,21 @@ class NumberStoppingCriteria(StoppingCriteria): if ( decoded.count(".") == 1 - and len(decoded.strip().split(".")[1]) > self.precision + and len(decoded.replace(" ", "").split(".")[1]) > self.precision + ): + return True + + if ( + len(decoded) > 1 + and "," in decoded + and any(c.isdigit() for c in decoded.split(",")[0]) ): return True if ( len(decoded) > 1 and any(c.isdigit() for c in decoded) - and decoded[-1] in [" ", "\n"] + and ("," in decoded or decoded[-1] in (" ", "\n")) ): return True @@ -71,9 +78,16 @@ class OutputNumbersTokens(LogitsWarper): for _, token_id in tokenizer.get_vocab().items(): token_str = tokenizer.decode(token_id).strip() - if token_str == "" or ( - all(c.isdigit() or c == "." for c in token_str) - and token_str.count(".") <= 1 + if ( + token_str == "" + or ( + all(c.isdigit() or c == "." for c in token_str) + and token_str.count(".") <= 1 + ) or ( + "," in token_str + and all(c.isdigit() or c == "." for c in token_str.split(",")[0]) + and token_str.count(".") <= 1 + ) ): self.allowed_mask[token_id] = True @@ -106,10 +120,17 @@ class IntegerStoppingCriteria(StoppingCriteria): if len(decoded.strip()) > self.max_digits: return True + if ( + len(decoded) > 1 + and "," in decoded + and any(c.isdigit() for c in decoded.split(",")[0]) + ): + return True + if ( len(decoded) > 1 and any(c.isdigit() for c in decoded) - and decoded[-1] in [" ", "\n"] + and decoded[-1] in (" ", "\n") ): return True @@ -125,7 +146,11 @@ class OutputIntegersTokens(LogitsWarper): 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): + if ( + token_str == "" + or all(c.isdigit() for c in token_str) + or "," in token_str and all(c.isdigit() for c in token_str.split(",")[0]) + ): self.allowed_mask[token_id] = True def __call__(self, _, scores): diff --git a/jsonformer/main.py b/jsonformer/main.py index b7bc5b4..f9e0280 100644 --- a/jsonformer/main.py +++ b/jsonformer/main.py @@ -76,7 +76,9 @@ class Jsonformer: response = self.tokenizer.decode(response[0], skip_special_tokens=True) response = response[len(prompt) :] - response = response.strip().rstrip(".") + if "," in response: + response = response.split(",")[0] + response = response.replace(" ", "").rstrip(".") self.debug("[generate_number]", response) try: return float(response) @@ -106,7 +108,9 @@ class Jsonformer: response = self.tokenizer.decode(response[0], skip_special_tokens=True) response = response[len(prompt) :] - response = response.strip() + if "," in response: + response = response.split(",")[0] + response = response.replace(" ", "") self.debug("[generate_integer]", response) try: return int(response)