From 77cead0e267d5e6aa29e80c96e55eaa9ace7dc64 Mon Sep 17 00:00:00 2001 From: wassname Date: Sat, 11 May 2024 10:54:51 +0800 Subject: [PATCH] range_mean --- prob_jsonformer/main.py | 22 ++++++++++++++++++++++ 1 file changed, 22 insertions(+) diff --git a/prob_jsonformer/main.py b/prob_jsonformer/main.py index 34d1d9b..538febb 100644 --- a/prob_jsonformer/main.py +++ b/prob_jsonformer/main.py @@ -182,6 +182,9 @@ class Jsonformer: return response.split('"')[0].strip() def generate_choice_probs(self, choices) -> str: + """ + This is not in the json schema, but can be usefull for effeciently getting the prob distibution over choices + """ prompt = self.get_prompt() + '"' self.debug("[generate_string_prob]", prompt, is_prompt=True) input_ids = self.tokenizer.encode(prompt, return_tensors="pt").to( @@ -193,6 +196,19 @@ class Jsonformer: r = list(choice_tree(self.model, self.tokenizer, input_ids, choices_tokens)) return r # json.dumps(r) + def generate_range_mean(self, range_min: float, range_max: float) -> float: + """ + This is not in the json schema, but can be usefull for effeciently generating the weighted mean from a range of integers + """ + choices = [str(n) for n in range(int(range_min), int(range_max) + 1)] + result = self.generate_choice_probs(choices) + + # now do a weighted average + total = 0.0 + for r in result: + total += float(r["choice"]) * r["prob"] + return total + def generate_enum(self, enum_values: Set[str]) -> str: prompt = self.get_prompt() self.debug("[generate_enum]", prompt, is_prompt=True) @@ -322,6 +338,12 @@ class Jsonformer: else: obj.append(self.generation_marker) return self.generate_choice_probs(schema["enum"]) + elif schema_type == "range_mean": + if key: + obj[key] = self.generation_marker + else: + obj.append(self.generation_marker) + return self.generate_range_mean(schema["minimum"], schema["maximum"]) elif schema_type == "enum": if key: obj[key] = self.generation_marker