| |
|
|
| |
| |
|
|
| import re |
| import numpy as np |
|
|
|
|
| def extract_solution(solution_str): |
| answer_pattern = r"<answer>(.*?)</answer>" |
| matches = re.findall(answer_pattern, solution_str, re.DOTALL) |
| if matches: |
| final_answer = matches[-1].strip() |
| else: |
| final_answer = None |
| return final_answer |
|
|
|
|
| def validate_equation(equation_str, available_numbers): |
| """Validate that equation only uses available numbers and each number once.""" |
| try: |
| |
| numbers_in_eq = [int(n) for n in re.findall(r"\d+", equation_str)] |
|
|
| |
| available_numbers = sorted(available_numbers) |
| numbers_in_eq = sorted(numbers_in_eq) |
|
|
| |
| return numbers_in_eq == available_numbers |
| except: |
| return False |
|
|
|
|
| def evaluate_equation(equation_str): |
| """Safely evaluate the arithmetic equation using eval() with precautions.""" |
| try: |
| |
| allowed_pattern = r"^[\d+\-*/().\s]+$" |
| if not re.match(allowed_pattern, equation_str): |
| raise ValueError("Invalid characters in equation.") |
|
|
| |
| result = eval(equation_str, {"__builtins__": None}, {}) |
| return result |
| except Exception as e: |
| return None |
|
|
|
|
| def compute_score(solution_str, ground_truth, method="strict", format_score=0.1, score=1.0): |
| """The scoring function for countdown task. |
| |
| Args: |
| solution_str: the solution text |
| ground_truth: dictionary containing target number and available numbers |
| method: the method to extract the solution |
| format_score: the score for correct format but wrong answer |
| score: the score for the correct answer |
| """ |
| target = ground_truth["target"] |
| numbers = ground_truth["numbers"] |
|
|
| equation = extract_solution(solution_str=solution_str) |
| if "wait" in solution_str: |
| do_print = True |
| do_print = np.random.rand() < 0.4 |
| if do_print: |
| print(f"--------------------------------") |
| print(f"Target: {target} | Numbers: {numbers}") |
| print(f"Extracted equation: {equation}") |
| print(f"Solution string: {solution_str}") |
|
|
| if equation is None: |
| if do_print: |
| print(f"No equation found") |
| return 0 |
|
|
| |
| if not validate_equation(equation, numbers): |
| if do_print: |
| print(f"Invalid equation") |
| return format_score |
|
|
| |
| try: |
| result = evaluate_equation(equation) |
| if result is None: |
| if do_print: |
| print(f"Could not evaluate equation") |
| return format_score |
|
|
| if abs(result - target) < 1e-5: |
| if do_print: |
| print(f"Correct equation: {equation} = {result}") |
| return score |
| else: |
| if do_print: |
| print(f"Wrong result: equation = {result}, target = {target}") |
| return format_score |
| except: |
| if do_print: |
| print(f"Error evaluating equation") |
| return format_score |
|
|
|
|
| class Parser: |
| @classmethod |
| def extract_answer_gsm8k(cls, generated_text): |
| """Extract the first numerical answer following '####' in the generated text.""" |
| try: |
| |
| |
| match = re.search(r"####\s*\$?([\d,]+(?:\.\d+)?)", generated_text) |
| if match: |
| return float(match.group(1).replace(",", "")) |
| except Exception as e: |
| print(f"Error extracting answer: {e}, Text: {generated_text[:100]}") |
| return None |
|
|
| @classmethod |
| def extract_answer_boxed(cls, generated_text): |
| """Extract the first numerical answer following '####' in the generated text.""" |
| try: |
| pred = remove_boxed(last_boxed_only_string(generated_text)) |
| except: |
| pred = generated_text |
| return pred |
|
|
| @classmethod |
| def extract_answer_boxed_ctd(cls, generated_text): |
| """Extract the first numerical answer following '####' in the generated text.""" |
| pred = Parser.extract_answer_boxed(generated_text) |
| pred = pred.replace(r"\div", "/").replace("\times", "*").replace(r"\cdot", "*") |
| return pred |
|
|
| @classmethod |
| def extract_answer_grpo_ctd(cls, generated_text): |
| """Extract the first numerical answer following '####' in the generated text.""" |
| pred = extract_solution(generated_text) |
| print(generated_text) |
| print(pred) |
| if pred is not None: |
|
|
| pred = pred.replace(r"\div", "/").replace("\times", "*").replace(r"\cdot", "*") |
|
|
| return pred |
|
|
| @classmethod |
| def extract_answer_sudoku(cls, solution_str): |
| """Extract the Sudoku solution from the generated text.""" |
| answer_pattern = r"<answer>(.*?)</answer>" |
| matches = re.findall(answer_pattern, solution_str, re.DOTALL) |
| if matches: |
| |
| final_answer = re.sub(r"\s", "", matches[-1].strip()) |
| return final_answer |
| return None |
|
|
|
|
| def is_equiv(str1, str2, verbose=False): |
| if type(str1) == float or type(str2) == float: |
| try: |
| return abs(float(str1) - float(str2)) < 1e-6 |
| except: |
| return False |
| if str1 is None and str2 is None: |
| print("WARNING: Both None") |
| return True |
| if str1 is None or str2 is None: |
| return False |
|
|
| try: |
| ss1 = strip_string(str1) |
| ss2 = strip_string(str2) |
| if verbose: |
| print(ss1, ss2) |
| return ss1 == ss2 |
| except Exception: |
| return str1 == str2 |
|
|
|
|
| def remove_boxed(s): |
| if "\\boxed " in s: |
| left = "\\boxed " |
| assert s[: len(left)] == left |
| return s[len(left) :] |
|
|
| left = "\\boxed{" |
|
|
| try: |
| assert s[: len(left)] == left |
| assert s[-1] == "}" |
|
|
| return s[len(left) : -1] |
| except: |
| return s |
|
|
|
|
| def last_boxed_only_string(string): |
| idx = string.rfind("\\boxed") |
| if "\\boxed " in string: |
| return "\\boxed " + string.split("\\boxed ")[-1].split("$")[0] |
| if idx < 0: |
| idx = string.rfind("\\fbox") |
| if idx < 0: |
| return string |
|
|
| i = idx |
| right_brace_idx = None |
| num_left_braces_open = 0 |
| while i < len(string): |
| if string[i] == "{": |
| num_left_braces_open += 1 |
| if string[i] == "}": |
| num_left_braces_open -= 1 |
| if num_left_braces_open == 0: |
| right_brace_idx = i |
| break |
| i += 1 |
|
|
| if right_brace_idx is None: |
| retval = None |
| else: |
| retval = string[idx : right_brace_idx + 1] |
|
|
| return retval |
|
|
|
|
| def fix_fracs(string): |
| substrs = string.split("\\frac") |
| new_str = substrs[0] |
| if len(substrs) > 1: |
| substrs = substrs[1:] |
| for substr in substrs: |
| new_str += "\\frac" |
| if substr[0] == "{": |
| new_str += substr |
| else: |
| try: |
| assert len(substr) >= 2 |
| except AssertionError: |
| return string |
| a = substr[0] |
| b = substr[1] |
| if b != "{": |
| if len(substr) > 2: |
| post_substr = substr[2:] |
| new_str += "{" + a + "}{" + b + "}" + post_substr |
| else: |
| new_str += "{" + a + "}{" + b + "}" |
| else: |
| if len(substr) > 2: |
| post_substr = substr[2:] |
| new_str += "{" + a + "}" + b + post_substr |
| else: |
| new_str += "{" + a + "}" + b |
| string = new_str |
| return string |
|
|
|
|
| def fix_a_slash_b(string): |
| if len(string.split("/")) != 2: |
| return string |
| a = string.split("/")[0] |
| b = string.split("/")[1] |
| try: |
| a = int(a) |
| b = int(b) |
| assert string == "{}/{}".format(a, b) |
| new_string = "\\frac{" + str(a) + "}{" + str(b) + "}" |
| return new_string |
| except AssertionError: |
| return string |
|
|
|
|
| def remove_right_units(string): |
| |
| if "\\text{ " in string: |
| splits = string.split("\\text{ ") |
| assert len(splits) == 2 |
| return splits[0] |
| else: |
| return string |
|
|
|
|
| def fix_sqrt(string): |
| if "\\sqrt" not in string: |
| return string |
| splits = string.split("\\sqrt") |
| new_string = splits[0] |
| for split in splits[1:]: |
| if split[0] != "{": |
| a = split[0] |
| new_substr = "\\sqrt{" + a + "}" + split[1:] |
| else: |
| new_substr = "\\sqrt" + split |
| new_string += new_substr |
| return new_string |
|
|
|
|
| def strip_string(string): |
| |
| string = string.replace("\n", "") |
|
|
| |
| string = string.replace("\\!", "") |
|
|
| |
| string = string.replace("\\\\", "\\") |
|
|
| |
| string = string.replace("tfrac", "frac") |
| string = string.replace("dfrac", "frac") |
|
|
| |
| string = string.replace("\\left", "") |
| string = string.replace("\\right", "") |
|
|
| |
| string = string.replace("^{\\circ}", "") |
| string = string.replace("^\\circ", "") |
|
|
| |
| string = string.replace("\\$", "") |
|
|
| |
| string = remove_right_units(string) |
|
|
| |
| string = string.replace("\\%", "") |
| string = string.replace("\%", "") |
|
|
| |
| string = string.replace(" .", " 0.") |
| string = string.replace("{.", "{0.") |
| |
| if len(string) == 0: |
| return string |
| if string[0] == ".": |
| string = "0" + string |
|
|
| |
| if len(string.split("=")) == 2: |
| if len(string.split("=")[0]) <= 2: |
| string = string.split("=")[1] |
|
|
| |
| string = fix_sqrt(string) |
|
|
| |
| string = string.replace(" ", "") |
|
|
| |
| string = fix_fracs(string) |
|
|
| |
| if string == "0.5": |
| string = "\\frac{1}{2}" |
|
|
| |
| string = fix_a_slash_b(string) |
|
|
| return string |
|
|
|
|
| def validate_equation(equation_str, available_numbers): |
| """Validate that equation only uses available numbers and each number once.""" |
| try: |
| |
| numbers_in_eq = [int(n) for n in re.findall(r"\d+", equation_str)] |
|
|
| |
| available_numbers = sorted(available_numbers) |
| numbers_in_eq = sorted(numbers_in_eq) |
|
|
| |
| return numbers_in_eq == available_numbers |
| except: |
| return False |
|
|
|
|
| def evaluate_equation(equation_str): |
| """Safely evaluate the arithmetic equation using eval() with precautions.""" |
| try: |
| |
| allowed_pattern = r"^[\d+\-*/().\s]+$" |
| if not re.match(allowed_pattern, equation_str): |
| raise ValueError("Invalid characters in equation.") |
|
|
| |
| result = eval(equation_str.strip(), {"__builtins__": None}, {}) |
| return result |
| except Exception as e: |
| return float("Inf") |
|
|