[Bench] fix MMMU answer-extraction regex dropping multi-line responses (#23864)

This commit is contained in:
Xinyuan Tong
2026-04-29 14:48:49 +08:00
committed by GitHub
parent 2c41ef4c93
commit 832b4f59ed
2 changed files with 207 additions and 1 deletions
+24 -1
View File
@@ -38,7 +38,7 @@ class EvalArgs:
concurrency: int = 1 concurrency: int = 1
max_new_tokens: Optional[int] = None max_new_tokens: Optional[int] = None
temperature: Optional[float] = None temperature: Optional[float] = None
response_answer_regex: str = "(.*)" response_answer_regex: str = "(?s)(.*)"
lora_path: Optional[str] = None lora_path: Optional[str] = None
reasoning_effort: Optional[str] = None reasoning_effort: Optional[str] = None
@@ -273,11 +273,34 @@ def get_sampling_params(eval_args):
# ----------- Process Multi-choice ------------- # ----------- Process Multi-choice -------------
def _parse_explicit_multi_choice_answer(response, all_choices):
choice_map = {choice.upper(): choice for choice in all_choices}
matches = []
answer_pattern = r"\banswer\s*:\s*\*{0,2}\s*\(?([A-Z])\)?\s*\*{0,2}(?![A-Za-z])"
for match in re.finditer(answer_pattern, response, flags=re.IGNORECASE):
candidate = match.group(1).upper()
if candidate in choice_map:
matches.append((match.start(1), choice_map[candidate]))
final_letter_pattern = r"(?:^|\n)\s*\*{0,2}\s*\(?([A-Z])\)?\s*\*{0,2}\s*\.?\s*$"
for match in re.finditer(final_letter_pattern, response, flags=re.IGNORECASE):
candidate = match.group(1).upper()
if candidate in choice_map:
matches.append((match.start(1), choice_map[candidate]))
return max(matches)[1] if matches else None
def parse_multi_choice_response(response, all_choices, index2ans): def parse_multi_choice_response(response, all_choices, index2ans):
""" """
Parse the prediction from the generated response. Parse the prediction from the generated response.
Return the predicted index e.g., A, B, C, D. Return the predicted index e.g., A, B, C, D.
""" """
explicit_answer = _parse_explicit_multi_choice_answer(response, all_choices)
if explicit_answer is not None:
return explicit_answer
for char in [",", ".", "!", "?", ";", ":", "'"]: for char in [",", ".", "!", "?", ";", ":", "'"]:
response = response.strip(char) response = response.strip(char)
response = " " + response + " " # add space to avoid partial match response = " " + response + " " # add space to avoid partial match
@@ -0,0 +1,183 @@
import importlib.util
import re
import sys
import types
import unittest
from pathlib import Path
try:
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
except ModuleNotFoundError:
CustomTestCase = unittest.TestCase
def register_cpu_ci(*args, **kwargs):
pass
register_cpu_ci(est_time=5, suite="stage-a-test-cpu")
def _load_mmmu_eval_utils():
repo_root = Path(__file__).resolve().parents[4]
module_path = repo_root / "benchmark" / "mmmu" / "eval_utils.py"
module_name = "_test_mmmu_eval_utils"
stub_modules = {
"data_utils": _build_data_utils_stub(),
"datasets": _build_datasets_stub(),
"numpy": _build_numpy_stub(),
"torch": types.ModuleType("torch"),
"tqdm": _build_tqdm_stub(),
}
previous_modules = {name: sys.modules.get(name) for name in stub_modules}
sys.modules.update(stub_modules)
spec = importlib.util.spec_from_file_location(module_name, module_path)
module = importlib.util.module_from_spec(spec)
sys.modules[module_name] = module
try:
spec.loader.exec_module(module)
finally:
for name, previous_module in previous_modules.items():
if previous_module is None:
sys.modules.pop(name, None)
else:
sys.modules[name] = previous_module
return module
def _build_data_utils_stub():
module = types.ModuleType("data_utils")
module.CAT_SHORT2LONG = {}
module.DOMAIN_CAT2SUB_CAT = {}
def _unused(*args, **kwargs):
raise AssertionError("Unexpected data_utils call in MMMU parser unit test")
module.construct_prompt = _unused
module.load_yaml = _unused
module.process_single_sample = _unused
module.save_json = _unused
return module
def _build_datasets_stub():
module = types.ModuleType("datasets")
def _unused(*args, **kwargs):
raise AssertionError("Unexpected datasets call in MMMU parser unit test")
module.concatenate_datasets = _unused
module.load_dataset = _unused
return module
def _build_numpy_stub():
module = types.ModuleType("numpy")
module.argmax = lambda values: max(range(len(values)), key=values.__getitem__)
return module
def _build_tqdm_stub():
module = types.ModuleType("tqdm")
module.tqdm = lambda iterable=None, *args, **kwargs: iterable
return module
class TestMMMUEvalUtils(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.eval_utils = _load_mmmu_eval_utils()
def test_default_response_answer_regex_captures_multiline_response(self):
response = "Based on the diagram, compare the labeled points.\nAnswer: B"
answer = re.search(self.eval_utils.EvalArgs.response_answer_regex, response)
self.assertIsNotNone(answer)
self.assertEqual(answer.group(1), response)
def test_default_regex_extraction_preserves_multiline_answer_for_processing(self):
response = "Based on the diagram, compare the labeled points.\nAnswer: B"
sample = self._multiple_choice_sample(response)
answer = re.search(self.eval_utils.EvalArgs.response_answer_regex, response)
answer_dict = {}
out_samples = {}
previous_random_choice = self.eval_utils.random.choice
self.eval_utils.random.choice = lambda choices: "A"
try:
self.eval_utils.process_result(
answer.group(1).strip() if answer else response,
sample,
answer_dict,
out_samples,
)
finally:
self.eval_utils.random.choice = previous_random_choice
self.assertEqual(out_samples["sample-1"]["pred_ans"], "B")
def test_parse_multi_choice_prefers_explicit_answer_marker_after_copied_options(
self,
):
response = (
"The options are:\n"
"(A) red\n"
"(B) blue\n"
"(C) green\n"
"(D) yellow\n"
"Answer: B"
)
pred_ans = self.eval_utils.parse_multi_choice_response(
response, ["A", "B", "C", "D"], self._index_to_answer()
)
self.assertEqual(pred_ans, "B")
def test_parse_multi_choice_prefers_final_standalone_letter_after_copied_options(
self,
):
response = (
"The options are:\n"
"(A) red\n"
"(B) blue\n"
"(C) green\n"
"(D) yellow\n"
"The diagram rules out the other labels.\n"
"**B**"
)
pred_ans = self.eval_utils.parse_multi_choice_response(
response, ["A", "B", "C", "D"], self._index_to_answer()
)
self.assertEqual(pred_ans, "B")
def test_parse_multi_choice_prefers_latest_explicit_answer(self):
response = "Initial thought: Answer: A\nAfter checking the image again:\n**B**"
pred_ans = self.eval_utils.parse_multi_choice_response(
response, ["A", "B", "C", "D"], self._index_to_answer()
)
self.assertEqual(pred_ans, "B")
def _multiple_choice_sample(self, response):
return {
"id": "sample-1",
"question_type": "multiple-choice",
"all_choices": ["A", "B", "C", "D"],
"index2ans": self._index_to_answer(),
"answer": "B",
"original_response": response,
}
def _index_to_answer(self):
return {"A": "red", "B": "blue", "C": "green", "D": "yellow"}
if __name__ == "__main__":
unittest.main()