| import asyncio |
| import re |
| from itertools import islice, zip_longest |
|
|
| from sympy.parsing.latex import parse_latex |
|
|
| try: |
| from math_verify import parse, verify |
| except ImportError: |
| print("math_verify is not installed in this environment") |
| parse = None |
| verify = None |
|
|
|
|
| def repeatness(s: str): |
| def ranks(l): |
| index = {v: i for i, v in enumerate(sorted(set(l)))} |
| return [index[v] for v in l] |
|
|
| def suffixArray(s): |
| line = ranks(s) |
| n, k, ans, sa = len(s), 1, line, [0] * len(s) |
| while k < n - 1: |
| line = ranks(list(zip_longest(line, islice(line, k, None), fillvalue=-1))) |
| ans, k = line, k << 1 |
| for i, k in enumerate(ans): |
| sa[k] = i |
| return ans, sa |
|
|
| def lcp(arr, suffixArr, inv_suff): |
| n, ans, k = len(arr), [0] * len(arr), 0 |
|
|
| for i in range(n): |
| if inv_suff[i] == n - 1: |
| k = 0 |
| continue |
|
|
| j = suffixArr[inv_suff[i] + 1] |
| while i + k < n and j + k < n and arr[i + k] == arr[j + k]: |
| k += 1 |
|
|
| ans[inv_suff[i]] = k |
| if k > 0: |
| k -= 1 |
|
|
| return ans |
|
|
| arr = [ord(i) for i in s] |
| n = len(arr) |
| if n <= 1: |
| return 0 |
| c, sa = suffixArray(arr) |
| cnt = sum(lcp(arr, sa, c)) |
|
|
| return (cnt * 2 / (n * (n + 1))) > 0.2 |
|
|
|
|
| SUBSTITUTIONS = [ |
| ("an ", ""), |
| ("a ", ""), |
| (".$", "$"), |
| ("\\$", ""), |
| (r"\ ", ""), |
| (" ", ""), |
| ("mbox", "text"), |
| (",\\text{and}", ","), |
| ("\\text{and}", ","), |
| ("\\text{m}", "\\text{}"), |
| ] |
|
|
|
|
| REMOVED_EXPRESSIONS = [ |
| "square", |
| "ways", |
| "integers", |
| "dollars", |
| "mph", |
| "inches", |
| "ft", |
| "hours", |
| "km", |
| "units", |
| "\\ldots", |
| "sue", |
| "points", |
| "feet", |
| "minutes", |
| "digits", |
| "cents", |
| "degrees", |
| "cm", |
| "gm", |
| "pounds", |
| "meters", |
| "meals", |
| "edges", |
| "students", |
| "childrentickets", |
| "multiples", |
| "\\text{s}", |
| "\\text{.}", |
| "\\text{\ns}", |
| "\\text{}^2", |
| "\\text{}^3", |
| "\\text{\n}", |
| "\\text{}", |
| r"\mathrm{th}", |
| r"^\circ", |
| r"^{\circ}", |
| r"\;", |
| r",\!", |
| "{,}", |
| '"', |
| "\\dots", |
| ] |
|
|
|
|
| def normalize_final_answer(final_answer: str) -> str: |
| """ |
| Normalize a final answer to a quantitative reasoning question. |
| This code comes from https://arxiv.org/pdf/2206.14858.pdf, page18. |
| """ |
| |
|
|
| for before, after in SUBSTITUTIONS: |
| final_answer = final_answer.replace(before, after) |
| for expr in REMOVED_EXPRESSIONS: |
| final_answer = final_answer.replace(expr, "") |
|
|
| |
| |
| final_answer = re.sub(r"(.*?)(\$)(.*?)(\$)(.*)", "$\\3$", final_answer) |
| final_answer = re.sub(r"(\\text\{)(.*?)(\})", "\\2", final_answer) |
| final_answer = re.sub(r"(\\textbf\{)(.*?)(\})", "\\2", final_answer) |
| final_answer = re.sub(r"(\\overline\{)(.*?)(\})", "\\2", final_answer) |
| final_answer = re.sub(r"(\\boxed\{)(.*)(\})", "\\2", final_answer) |
|
|
| |
| |
| |
| |
| |
| |
| final_answer = re.sub(r"(frac)([^{])(.)", "frac{\\2}{\\3}", final_answer) |
| final_answer = re.sub(r"(sqrt)([^{])", "sqrt{\\2}", final_answer) |
| final_answer = final_answer.replace("$", "") |
|
|
| |
| if final_answer.replace(",", "").isdigit(): |
| final_answer = final_answer.replace(",", "") |
|
|
| return final_answer |
|
|
|
|
| def latex_eval(latex): |
| sym = parse_latex(latex) |
| val = sym.evalf() |
| return sym, val |
|
|
|
|
| def _is_latex_equal(str1, str2): |
| try: |
| sym1, val1 = latex_eval(str1) |
| sym2, val2 = latex_eval(str2) |
| if sym1 == sym2 or val1 == val2: |
| return True |
| else: |
| raise ValueError |
| except Exception: |
| try: |
| norm1, norm2 = normalize_final_answer(str1), normalize_final_answer(str2) |
| sym1, val1 = latex_eval(norm1) |
| sym2, val2 = latex_eval(norm2) |
| if sym1 == sym2 or val1 == val2: |
| return True |
| except Exception: |
| return norm1 == norm2 |
| return False |
|
|
|
|
| async def is_latex_equal(str1, str2, executor, math_mode="legacy"): |
| if math_mode == "legacy": |
| if (len(str1) > 128 and repeatness(str1)) or (len(str2) > 128 and repeatness(str2)): |
| return False |
|
|
| try: |
| loop = asyncio.get_event_loop() |
| task = loop.run_in_executor(executor, _is_latex_equal, str1, str2) |
| result = await asyncio.wait_for(task, timeout=1.0) |
| return result |
| except asyncio.exceptions.TimeoutError: |
| return False |
| elif math_mode == "math_verify": |
| try: |
| loop = asyncio.get_event_loop() |
| task = loop.run_in_executor(executor, verify, parse(str1), parse(str2)) |
| result = await asyncio.wait_for(task, timeout=1.0) |
| return result |
| except asyncio.exceptions.TimeoutError: |
| return False |
| else: |
| raise NotImplementedError(f"Math mode {math_mode} is not implemented") |
|
|
|
|
| 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 Exception: |
| 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 Exception: |
| 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 = string.replace("$", "") |
| 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 is_equiv(str1, str2, verbose=False) -> bool: |
| 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) |
| try: |
| return float(ss1) == (float(ss2)) |
| except Exception: |
| return ss1 == ss2 |
| except Exception: |
| return str1 == str2 |
|
|
|
|
| def last_boxed_only_string(string): |
| idx = string.rfind("\\boxed") |
| if idx < 0: |
| idx = string.rfind("\\fbox") |
| if idx < 0: |
| return None |
|
|
| 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 remove_boxed(s): |
| left = "\\boxed{" |
| try: |
| assert s[: len(left)] == left |
| assert s[-1] == "}" |
| return s[len(left) : -1] |
| except Exception: |
| return None |
|
|
|
|
| def get_answer_str(s: str) -> str: |
| res = remove_boxed(last_boxed_only_string(s)) |
| if res is not None: |
| return res |
| return s |
|
|
|
|
| async def is_equal(str1, str2, executor, math_mode="legacy"): |
| first_equal = is_equiv(str1, str2) |
| if first_equal: |
| return True |
| return await is_latex_equal(str1, str2, executor, math_mode) |
|
|
|
|
| def solution2answer(solution: str, math_mode="eval_peeking") -> str: |
| answer = solution |
| if math_mode == "eval_peeking": |
| answer = get_answer_str(solution) |
| else: |
| raise ValueError(f"Invalid math_mode: {math_mode}") |
| return answer |
|
|
|
|
| def get_final_answer(output: str) -> str: |
| output = output.replace("is:", "is").replace("answer:", "answer is").strip() |
| if output.endswith("."): |
| output = output[:-1] |
| if ".$" in output: |
| output = output.replace(".$", "$") |
| pattern_list = [ |
| r"answer is (-?\d+\.?\d*)$", |
| r"answer is (.+?)$", |
| ] |
| matches = [] |
| for pat in pattern_list: |
| matches = re.findall(pat, output, re.S) |
| if matches: |
| return get_answer_str(matches[0]) |
|
|
| return get_answer_str(output) |