Repository navigation
Expand file tree
/
Copy pathmath_answer_extractor.py
More file actions
145 lines (126 loc) · 3.95 KB
/
Copy pathmath_answer_extractor.py
File metadata and controls
145 lines (126 loc) · 3.95 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
from sympy.parsing import sympy_parser as spp
from sympy.core.sympify import SympifyError
from tokenize import TokenError
import re
RE_NUMBER = re.compile(r"-?(?:\d+/\d+|\d+(?:\.\d+)?(?:[eE][+-]?\d+)?)")
LATEX_FIXES = [
(r"\\left\s*", ""),
(r"\\right\s*", ""),
(r"\\,|\\!|\\;|\\:", ""),
(r"\\cdot", "*"),
(r"\u00B7|\u00D7", "*"),
(r"\\\^\\circ", ""),
(r"\\dfrac", r"\\frac"),
(r"\\tfrac", r"\\frac"),
(r"°", ""),
]
RE_SPECIAL = re.compile(r"<\|[^>]+?\|>")
# extracting the final answer box from the generated answer
def extract_final_answer_box(answer: str) -> str | None:
box_start_idx = answer.rfind(r"\boxed")
if box_start_idx == -1:
return None
current_idx = box_start_idx + len("r\boxed")
while current_idx < len(answer) and answer[current_idx].isspace():
current_idx += 1
if current_idx == len(answer) or answer[current_idx] != "{":
return None
current_idx += 1
brace_depth = 1 #'we opened {'
content_start_idx = current_idx
while current_idx < len(answer) and brace_depth > 0:
if answer[current_idx] == "{":
brace_depth += 1
elif answer[current_idx] == "}":
brace_depth -= 1
current_idx += 1
if brace_depth != 0:
return None
return answer[content_start_idx : current_idx - 1]
def extract_final_answer(answer, fallback="number_then_full"):
box = extract_final_answer_box(answer)
result = box or ""
if result:
return result
if fallback:
numbers = RE_NUMBER.findall(answer)
if numbers:
result = numbers[-1]
elif fallback == "number_then_full":
result = answer
return result
def normalize_text(text):
if not text:
return ""
text = RE_SPECIAL.sub("", text).strip()
text = re.sub(r"\^\s*\{\s*\\circ\s*\}", "", text)
text = re.sub(r"\^\s*\\circ", "", text)
text = text.replace("°", "")
match = re.match(r"^\\text\{(?P<x>.+?)\}$", text)
if match:
text = match.group("x")
text = re.sub(r"\\\(|\\\)|\\\[|\\\]", "", text)
for pat, rep in LATEX_FIXES:
text = re.sub(pat, rep, text)
text = text.replace("\\%", "%").replace("$", "").replace("%", "")
text = re.sub(
r"\\sqrt\s*\{([^}]*)\}",
lambda match: f"sqrt({match.group(1)})",
text,
)
text = re.sub(
r"\\sqrt\s+([^\\\s{}]+)",
lambda match: f"sqrt({match.group(1)})",
text,
)
text = re.sub(
r"\\frac\s*\{([^{}]+)\}\s*\{([^{}]+)\}",
lambda match: f"({match.group(1)})/({match.group(2)})",
text,
)
text = re.sub(
r"\\frac\s+([^\s{}]+)\s+([^\s{}]+)",
lambda match: f"({match.group(1)})/({match.group(2)})",
text,
)
text = text.replace("^", "**")
text = re.sub(
r"(?<=\d)\s+(\d+/\d+)",
lambda match: "+" + match.group(1),
text,
)
text = re.sub(
r"(?<=\d),(?=\d\d\d(\D|$))",
"",
text,
)
return text.replace("{", "").replace("}", "").strip().lower()
# parser number: from string number to number
def sympy_parser(expr):
try:
return spp.parse_expr(
expr,
transformations=(
*spp.standard_transformations,
spp.implicit_multiplication_application,
),
evaluate=True,
)
except (SympifyError, SyntaxError, TypeError, IndexError, TokenError):
return None
def split_into_parts(exp):
if exp[0] in "[(" and exp[-1] in ")]" and "," in exp[1:-1]:
exp = exp[1:-1].split(",")
return [x.strip() for x in exp]
return [exp] if exp else []
def answer_verifier(prediction, truth):
if prediction == truth:
return True
prediction, truth = sympy_parser(prediction), sympy_parser(truth)
if prediction and truth:
try:
if prediction - truth == 0:
return True
except Exception:
pass
return False