diff --git a/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluate_utils/evaluate_dicts.py b/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluate_utils/evaluate_dicts.py index 76136c06a..b5778dc11 100644 --- a/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluate_utils/evaluate_dicts.py +++ b/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluate_utils/evaluate_dicts.py @@ -2,7 +2,7 @@ import numpy as np -from .utils import _align_bags +from .utils import _align_bags, _fix_comma def calculate_f1_score(precision, recall): @@ -43,7 +43,7 @@ def fix_number(number): " ".join(" ".join(copy_ans.split("$")).split("%")).split("sqft") ).strip() copy_ans = copy_ans.strip() - copy_ans = copy_ans.replace(",", ".") + copy_ans = _fix_comma(copy_ans) try: return float(copy_ans) except: diff --git a/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluate_utils/utils.py b/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluate_utils/utils.py index 82be9e4ad..dec157701 100644 --- a/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluate_utils/utils.py +++ b/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluate_utils/utils.py @@ -1,3 +1,4 @@ +import re from typing import Callable, List, Set import numpy as np @@ -23,3 +24,23 @@ def _align_bags( for row, column in zip(row_ind, col_ind): max_scores[row] = max(max_scores[row], scores[row, column]) return max_scores + + +def _fix_comma(number: str) -> str: + """ + Distinguishes US thousands separators from European decimal commas. + + Multiple commas, or a single comma followed by exactly 3 digits, are + treated as thousands separators and removed; a single comma followed by + 1 or 2 digits is treated as a decimal comma and replaced by a period. + Anything else is left untouched. + """ + if number.count(",") > 1: + return number.replace(",", "") + match = re.fullmatch(r"(-?\d+),(\d+)", number) + if match: + if len(match.group(2)) == 3: + return number.replace(",", "") + if len(match.group(2)) <= 2: + return number.replace(",", ".") + return number diff --git a/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluator.py b/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluator.py index 7910eadf8..3eb1e1dd9 100644 --- a/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluator.py +++ b/browsergym/assistantbench/src/browsergym/assistantbench/evaluation/evaluator.py @@ -5,6 +5,7 @@ import numpy as np from .evaluate_utils.evaluate_factory import get_evaluator +from .evaluate_utils.utils import _fix_comma def find_isnan(samp): @@ -60,7 +61,8 @@ def fix_number(number): " ".join(" ".join(copy_ans.split("$")).split("%")).split("sqft") ).strip() copy_ans = copy_ans.strip() - copy_ans = copy_ans.replace(",", ".").replace(" square kilometers", "") + copy_ans = copy_ans.replace(" square kilometers", "") + copy_ans = _fix_comma(copy_ans) try: return float(copy_ans), True except: diff --git a/tests/assistantbench/test_evaluation.py b/tests/assistantbench/test_evaluation.py index 4973d7158..cb3a829b8 100644 --- a/tests/assistantbench/test_evaluation.py +++ b/tests/assistantbench/test_evaluation.py @@ -52,6 +52,36 @@ def test_evaluate(original_id: str): assert has_ans == expected_has_ans +@pytest.mark.parametrize( + "prediction, gold_answer, expected_score", + [ + # US thousands separators must be removed before parsing (#404) + ("3,080,000", "3080000", 1.0), + ("1,000", "1000", 1.0), + # European decimal commas must keep scoring 1.0 + ("14,2", "14.2", 1.0), + # 1000x errors must no longer score 1.0 (#404) + ("1,000", "1", 0.0), + ("1,010", "1.01", 0.0), + ], +) +def test_evaluate_comma_numbers(prediction, gold_answer, expected_score): + + score, _ = question_scorer(prediction, gold_answer) + + assert score == expected_score + + +def test_evaluate_comma_numbers_in_dicts(): + + prediction = '[{"sender": "DHL", "price": "1,000"}]' + gold_answer = '{"sender": "DHL", "price": 1000}' + + score, _ = question_scorer(prediction, gold_answer) + + assert score == 1.0 + + @pytest.mark.parametrize( "original_id", [id for id in data_points.keys() if isinstance(data_points[id]["answer"], (str, float, int))],