Fix TestGLM41VPPAccuracy test flakiness (#14848)
This commit is contained in:
@@ -115,7 +115,11 @@ def run_eval(args):
|
|||||||
# VLM MMMU evaluation with fixed 100 examples by default
|
# VLM MMMU evaluation with fixed 100 examples by default
|
||||||
from sglang.test.simple_eval_mmmu_vlm import MMMUVLMEval
|
from sglang.test.simple_eval_mmmu_vlm import MMMUVLMEval
|
||||||
|
|
||||||
eval_obj = MMMUVLMEval(args.num_examples, args.num_threads)
|
eval_obj = MMMUVLMEval(
|
||||||
|
args.num_examples,
|
||||||
|
args.num_threads,
|
||||||
|
response_answer_regex=getattr(args, "response_answer_regex", None),
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Invalid eval name: {args.eval_name}")
|
raise ValueError(f"Invalid eval name: {args.eval_name}")
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import base64
|
import base64
|
||||||
import io
|
import io
|
||||||
|
import re
|
||||||
from typing import List, Optional, Tuple
|
from typing import List, Optional, Tuple
|
||||||
|
|
||||||
from datasets import concatenate_datasets, load_dataset
|
from datasets import concatenate_datasets, load_dataset
|
||||||
@@ -53,7 +54,11 @@ class MMMUVLMEval(Eval):
|
|||||||
}
|
}
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self, num_examples: Optional[int] = 100, num_threads: int = 32, seed: int = 42
|
self,
|
||||||
|
num_examples: Optional[int] = 100,
|
||||||
|
num_threads: int = 32,
|
||||||
|
seed: int = 42,
|
||||||
|
response_answer_regex: str = None,
|
||||||
):
|
):
|
||||||
"""Create MMMU VLM eval (Math subset, 100 fixed samples by default)."""
|
"""Create MMMU VLM eval (Math subset, 100 fixed samples by default)."""
|
||||||
self.num_examples = num_examples
|
self.num_examples = num_examples
|
||||||
@@ -61,6 +66,10 @@ class MMMUVLMEval(Eval):
|
|||||||
self.seed = seed
|
self.seed = seed
|
||||||
# Prepare samples deterministically across all MMMU subjects (validation split)
|
# Prepare samples deterministically across all MMMU subjects (validation split)
|
||||||
self.samples = self._prepare_mmmu_samples(self.num_examples)
|
self.samples = self._prepare_mmmu_samples(self.num_examples)
|
||||||
|
# For example, "<\|begin_of_box\|>foo<\|end_of_box\|>" could be used to extract "foo" as the answer from the response text
|
||||||
|
self.response_answer_regex = (
|
||||||
|
response_answer_regex if response_answer_regex is not None else "(.*)"
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _to_data_uri(image: Image.Image) -> str:
|
def _to_data_uri(image: Image.Image) -> str:
|
||||||
@@ -205,6 +214,14 @@ class MMMUVLMEval(Eval):
|
|||||||
# Sample
|
# Sample
|
||||||
response_text = sampler(prompt_messages)
|
response_text = sampler(prompt_messages)
|
||||||
response_text = response_text or ""
|
response_text = response_text or ""
|
||||||
|
match = (
|
||||||
|
re.search(self.response_answer_regex, response_text)
|
||||||
|
if response_text is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
response_text = (
|
||||||
|
match.group(1).strip() if match is not None else response_text
|
||||||
|
)
|
||||||
|
|
||||||
# Parse and score
|
# Parse and score
|
||||||
gold = sample["answer"]
|
gold = sample["answer"]
|
||||||
|
|||||||
@@ -339,6 +339,8 @@ class TestGLM41VPPAccuracy(unittest.TestCase):
|
|||||||
"--chunked-prefill-size",
|
"--chunked-prefill-size",
|
||||||
8192,
|
8192,
|
||||||
"--enable-multimodal",
|
"--enable-multimodal",
|
||||||
|
"--reasoning-parser",
|
||||||
|
"glm45",
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -353,10 +355,12 @@ class TestGLM41VPPAccuracy(unittest.TestCase):
|
|||||||
eval_name="mmmu",
|
eval_name="mmmu",
|
||||||
num_examples=None,
|
num_examples=None,
|
||||||
num_threads=32,
|
num_threads=32,
|
||||||
|
response_answer_regex="<\|begin_of_box\|>(.*)<\|end_of_box\|>",
|
||||||
)
|
)
|
||||||
|
|
||||||
metrics = run_eval(args)
|
metrics = run_eval(args)
|
||||||
print(f"{metrics=}")
|
print(f"{metrics=}")
|
||||||
self.assertGreater(metrics["score"], 0.55)
|
self.assertGreater(metrics["score"], 0.45)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
Reference in New Issue
Block a user