From 26ee0c08fc15daadba718c1aeaf1a0767f385e69 Mon Sep 17 00:00:00 2001 From: posiounus Date: Sat, 10 Jun 2023 17:29:56 -0400 Subject: [PATCH] fix bugs in choose_type_to_generate --- jsonformer/main.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/jsonformer/main.py b/jsonformer/main.py index f27558f..cebafe0 100644 --- a/jsonformer/main.py +++ b/jsonformer/main.py @@ -153,7 +153,7 @@ class Jsonformer: def choose_type_to_generate(self, possible_types: List[str]) -> str: possible_types = list(set(possible_types)) # remove duplicates - self.debug("[choosing type to generate]", possible_types) + self.debug("[choose_type_to_generate]", possible_types) if len(possible_types) < 1: raise ValueError(f"Union type must not be empty") elif len(possible_types) == 1: @@ -165,11 +165,11 @@ class Jsonformer: logits = output.logits[0, -1] max_type = None - max_logit = -1 + max_logit = -float("inf") for possible_type in possible_types: try: prefix_tokens = self.type_prefix_tokens[possible_type] - except IndexError: + except KeyError: raise ValueError(f"Unsupported schema type: {possible_type}") max_type_logit = logits[prefix_tokens].max() if max_type_logit > max_logit: @@ -178,7 +178,7 @@ class Jsonformer: if max_type is None: raise Exception("Unable to find best type to generate for union type") - self.debug("[chose type to generate]", max_type) + self.debug("[choose_type_to_generate]", max_type) return max_type def generate_value(