- perf_stats aggregator lives in eval/, not model/: the import failed silently and EVERY perf column was empty (not just ttft). Now warns on stderr instead of swallowing. - repeats > 1 get their own checkpoint key (:rep2, :rep3, ...): repeat 2 previously restored repeat 1's predictions and finished instantly with identical scores. rep1 keeps the legacy key (existing checkpoints still resume). - repeats summary: report the MEAN score and aggregate time/tokens over ALL runs (was: last run only). - README: six-benchmark command as the primary example. Co-Authored-By: Claude <noreply@anthropic.com>
146 lines
4.2 KiB
Python
146 lines
4.2 KiB
Python
"""Optional sympy-based math grader (adapted from OpenAI PRM800K, MIT license).
|
||
|
||
Imported lazily by the math_equal scorer; only needed when normalized string
|
||
equality is not enough (unreduced fractions, symbolic forms, units).
|
||
Requires: pip install sympy pylatexenc
|
||
"""
|
||
|
||
import logging
|
||
from typing import Optional
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
try:
|
||
import sympy
|
||
from pylatexenc import latex2text
|
||
from sympy.parsing import sympy_parser
|
||
|
||
_OK = True
|
||
except ImportError: # pragma: no cover - optional dependency
|
||
_OK = False
|
||
|
||
BAD_SUBSTRINGS = ('^{', '^(')
|
||
TUPLE_CHARS = '()[]'
|
||
|
||
|
||
def _sympy_parse(expr: str):
|
||
return sympy_parser.parse_expr(
|
||
expr.replace('^', '**'),
|
||
transformations=(
|
||
sympy_parser.standard_transformations
|
||
+ (sympy_parser.implicit_multiplication_application,)
|
||
),
|
||
)
|
||
|
||
|
||
def _parse_latex(expr: str) -> str:
|
||
expr = expr.replace('\\tfrac', '\\frac').replace('\\dfrac', '\\frac')
|
||
expr = latex2text.LatexNodes2Text().latex_to_text(expr)
|
||
return (expr.replace('√', 'sqrt').replace('π', 'pi').replace('·', '*')
|
||
.replace('×', '*').strip())
|
||
|
||
|
||
def _strip_commas(expr: str) -> str:
|
||
import re
|
||
|
||
return re.sub(r'(\d),(\d\d\d)(?=\D|$)', r'\1\2', expr)
|
||
|
||
|
||
def _is_float(x) -> bool:
|
||
try:
|
||
float(x)
|
||
return True
|
||
except (TypeError, ValueError):
|
||
return False
|
||
|
||
|
||
def _str_is_int(x: str) -> bool:
|
||
try:
|
||
return abs(float(x) - int(round(float(x)))) <= 1e-7
|
||
except (TypeError, ValueError):
|
||
return False
|
||
|
||
|
||
def _is_frac(expr: str) -> bool:
|
||
import re
|
||
|
||
return bool(re.fullmatch(r'-?[0-9]+.?/0*[1-9][0-9]*.?', expr or ''))
|
||
|
||
|
||
def _normalize(expr: Optional[str]) -> Optional[str]:
|
||
import re
|
||
|
||
if expr is None:
|
||
return None
|
||
m = re.fullmatch(r'\\text\{(.+?)\}', expr)
|
||
if m:
|
||
expr = m.group(1)
|
||
expr = (expr.replace('\\%', '%').replace('\\$', '$').replace('$', '').replace('%', '')
|
||
.replace(' or ', ' , ').replace(' and ', ' , ')
|
||
.replace('million', '*10^6').replace('billion', '*10^9'))
|
||
for unit in ('degree', 'cm', 'meter', 'mile', 'second', 'minute', 'hour',
|
||
'day', 'week', 'month', 'year', 'foot', 'feet', 'inch', 'yard'):
|
||
expr = re.sub(unit + r'(es)?(s)? *(\^[0-9]+)?', '', expr)
|
||
expr = re.sub(r'\^ *\\circ', '', expr)
|
||
if len(expr) > 1 and expr[0] == '{' and expr[-1] == '}':
|
||
expr = expr[1:-1]
|
||
expr = _strip_commas(expr)
|
||
if _is_float(expr):
|
||
try:
|
||
f = float(expr)
|
||
if abs(f - round(f)) <= 1e-7:
|
||
expr = str(int(round(f)))
|
||
except (TypeError, ValueError):
|
||
pass
|
||
if '\\' in expr:
|
||
try:
|
||
expr = _parse_latex(expr)
|
||
except Exception:
|
||
pass
|
||
expr = re.sub(r'- *', '-', expr.replace(' ', ''))
|
||
expr = expr.replace('{', '').replace('}', '').lower()
|
||
if _str_is_int(expr):
|
||
expr = str(int(round(float(expr))))
|
||
return expr
|
||
|
||
|
||
def _split_tuple(expr: str):
|
||
expr = _strip_commas(expr)
|
||
if (len(expr) > 2 and expr[0] in TUPLE_CHARS and expr[-1] in TUPLE_CHARS
|
||
and all(c not in expr[1:-1] for c in TUPLE_CHARS)):
|
||
return [e.strip() for e in expr[1:-1].split(',')]
|
||
return [expr]
|
||
|
||
|
||
def _sympy_equal(a: str, b: str) -> bool:
|
||
try:
|
||
diff = _sympy_parse(f'({a})-({b})')
|
||
return sympy.simplify(diff) == 0
|
||
except Exception:
|
||
return False
|
||
|
||
|
||
def grade_answer(given_answer: str, ground_truth: str) -> bool:
|
||
"""True iff equal under normalization or sympy simplification."""
|
||
if not _OK:
|
||
raise ImportError('pip install sympy pylatexenc for symbolic math grading')
|
||
if given_answer is None:
|
||
return False
|
||
a, b = _normalize(given_answer), _normalize(ground_truth)
|
||
if a == b:
|
||
return True
|
||
if not a or not b:
|
||
return False
|
||
a_elems, b_elems = _split_tuple(a), _split_tuple(b)
|
||
if len(a_elems) != len(b_elems):
|
||
return False
|
||
for x, y in zip(a_elems, b_elems):
|
||
if _is_frac(x) and _is_frac(y):
|
||
if x != y:
|
||
return False
|
||
elif _str_is_int(x) != _str_is_int(y):
|
||
return False
|
||
elif not _sympy_equal(x, y):
|
||
return False
|
||
return True
|