diff --git a/pyproject.toml b/pyproject.toml index 02d4f00..e911c0d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -18,7 +18,7 @@ dependencies = [ ] [project.optional-dependencies] -dev = ["pytest>=8.0", "ruff>=0.5", "mypy>=1.10"] +dev = ["pytest>=8.0", "ruff>=0.5,<0.16", "mypy>=1.10"] [build-system] requires = ["hatchling"] diff --git a/src/trinity/adapters/drop.py b/src/trinity/adapters/drop.py index 1314803..07fcbd4 100644 --- a/src/trinity/adapters/drop.py +++ b/src/trinity/adapters/drop.py @@ -130,9 +130,12 @@ def _maybe_unbox(segment: str) -> str: return boxed if boxed is not None else "" -#: Surrounding punctuation stripped from a token, EXCLUDING the signs ``+``/``-`` — a -#: leading sign is part of a number's value, not wrapping noise. -_STRIP_EDGE = "".join(c for c in string.punctuation if c not in "+-") +#: Surrounding punctuation stripped from a token, EXCLUDING the signs ``+``/``-`` and +#: the decimal point ``.`` — a leading sign or decimal point is part of a number's +#: value, not wrapping noise. A genuinely-trailing ``.`` (sentence period) is handled +#: by the dedicated rstrip retry in :func:`_normalize_token`, which can tell it apart +#: from a value-bearing leading point; a blanket edge-strip cannot. +_STRIP_EDGE = "".join(c for c in string.punctuation if c not in "+-.") def _normalize_token(raw: str) -> str: @@ -147,20 +150,27 @@ def _normalize_token(raw: str) -> str: dropped the ``-`` and left commas to break ``float()``) did not deliver. A token that is ALREADY a number is recognised before any punctuation is - stripped: the edge-strip set includes ``.``, so a leading-decimal token like - ``".5"`` would otherwise lose its point and normalize to ``"5.0"`` — equal to - a gold ``"5"`` (false positive) and unequal to the value-identical gold - ``"0.5"`` (false negative). The official DROP ``_remove_punc`` tests - ``_is_number`` first and leaves numbers untouched for exactly this reason.""" + stripped: a leading-decimal token like ``".5"`` must not lose its point and + normalize to ``"5.0"`` — equal to a gold ``"5"`` (false positive) and unequal + to the value-identical gold ``"0.5"`` (false negative). The official DROP + ``_remove_punc`` tests ``_is_number`` first and leaves numbers untouched for + exactly this reason (issue #423). The float-first path alone only covers the + *bare* token: with ``.`` in the edge-strip set, wrapped forms like ``"$.5"`` + and ``".5."`` still lost the leading point on the second-chance path. So the + edge strip excludes ``.`` entirely, and a genuinely-trailing period (``".5."``, + ``"16.."``) is retried with an explicit ``rstrip(".")`` — right-side dots are + sentence punctuation, left-side dots are value.""" try: return str(float(raw.replace(",", ""))) except ValueError: pass core = raw.strip(_STRIP_EDGE) - try: - return str(float(core.replace(",", ""))) - except ValueError: - return _PUNCT.sub("", raw) + for cand in (core, core.rstrip(".")): + try: + return str(float(cand.replace(",", ""))) + except ValueError: + continue + return _PUNCT.sub("", raw) def _split_internal_hyphens(token: str) -> list[str]: diff --git a/src/trinity/orchestration/reward.py b/src/trinity/orchestration/reward.py index 7192ef8..d196a8d 100644 --- a/src/trinity/orchestration/reward.py +++ b/src/trinity/orchestration/reward.py @@ -847,8 +847,13 @@ def normalize_math_answer(ans: str | None) -> str: # standalone atomic operands adjacent to division — including a lone # ``sqrt(...)`` call produced from ``\sqrt{...}`` — without eating the # function-call parentheses themselves (``(sqrt(2))`` -> ``sqrt(2)``). + # Include #434's ``base^(...)`` spelling: the strip runs *after* + # ``_normalize_braced_exponents``, so ``x^{2}`` / ``2^{10}`` already look + # like ``x^(2)`` / ``2^(10)``. Without that branch, ``\frac{x^{2}}{2}`` + # stays ``(x^(2))/2`` while the slash form is ``x^(2)/2`` (exact-match + # false negative; only sympy saves it). s = re.sub( - r"(^|/)\((pi|\d+|sqrt\([^()]*\)|[a-z](?:\^\{[^{}]*\}|\^[0-9a-z])?)\)(?=/|$)", + r"(^|/)\((pi|\d+|sqrt\([^()]*\)|(?:\d+|[a-z])(?:\^\([^()]*\)|\^\{[^{}]*\}|\^[0-9a-z])?)\)(?=/|$)", r"\1\2", s, ) diff --git a/tests/test_reward_nested_frac_unwrap.py b/tests/test_reward_nested_frac_unwrap.py index ea04967..df9f720 100644 --- a/tests/test_reward_nested_frac_unwrap.py +++ b/tests/test_reward_nested_frac_unwrap.py @@ -18,5 +18,9 @@ def test_plain_frac_still_matches(): def test_frac_with_superscript_operand(): # Braced superscript in the numerator is the same [^{}]+ failure mode. + # Exact normalize must agree — do not rely on the sympy fallback. assert "\\frac" not in normalize_math_answer(r"\frac{x^{2}}{2}") + assert normalize_math_answer(r"\frac{x^{2}}{2}") == normalize_math_answer(r"x^{2}/2") assert math_equal(r"\frac{x^{2}}{2}", r"x^{2}/2") + assert normalize_math_answer(r"\frac{2^{10}}{2}") == normalize_math_answer(r"2^{10}/2") + assert math_equal(r"\frac{2^{10}}{2}", r"2^{10}/2")