Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import re
from typing import Callable, List, Set

import numpy as np
Expand All @@ -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
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand Down
30 changes: 30 additions & 0 deletions tests/assistantbench/test_evaluation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))],
Expand Down