diff --git a/gpt_oss/evals/__main__.py b/gpt_oss/evals/__main__.py index 40d56c12..83ff06f0 100644 --- a/gpt_oss/evals/__main__.py +++ b/gpt_oss/evals/__main__.py @@ -3,14 +3,15 @@ from datetime import datetime from . import report -from .basic_eval import BasicEval -from .gpqa_eval import GPQAEval from .aime_eval import AIME25Eval -from .healthbench_eval import HealthBenchEval +from .basic_eval import BasicEval from .chat_completions_sampler import ( OPENAI_SYSTEM_MESSAGE_API, ChatCompletionsSampler, ) +from .cli_utils import resolve_gpqa_debug_mode, resolve_n_repeats, resolve_num_examples +from .gpqa_eval import GPQAEval +from .healthbench_eval import HealthBenchEval from .responses_sampler import ResponsesSampler @@ -95,24 +96,24 @@ def main(): ) def get_evals(eval_name, debug_mode): - num_examples = ( - args.examples if args.examples is not None else (5 if debug_mode else None) - ) + num_examples = resolve_num_examples(args.examples, debug_mode, 5) + healthbench_num_examples = resolve_num_examples(args.examples, debug_mode, 10) + n_repeats = resolve_n_repeats(args.examples, debug_mode) # Set num_examples = None to reproduce full evals match eval_name: case "basic": return BasicEval() case "gpqa": return GPQAEval( - n_repeats=1 if args.debug else 8, + n_repeats=n_repeats, num_examples=num_examples, - debug=debug_mode, + debug=resolve_gpqa_debug_mode(args.examples, debug_mode), n_threads=args.n_threads or 1, ) case "healthbench": return HealthBenchEval( grader_model=grading_sampler, - num_examples=10 if debug_mode else num_examples, + num_examples=healthbench_num_examples, n_repeats=1, n_threads=args.n_threads or 1, subset_name=None, @@ -120,7 +121,7 @@ def get_evals(eval_name, debug_mode): case "healthbench_hard": return HealthBenchEval( grader_model=grading_sampler, - num_examples=10 if debug_mode else num_examples, + num_examples=healthbench_num_examples, n_repeats=1, n_threads=args.n_threads or 1, subset_name="hard", @@ -128,14 +129,14 @@ def get_evals(eval_name, debug_mode): case "healthbench_consensus": return HealthBenchEval( grader_model=grading_sampler, - num_examples=10 if debug_mode else num_examples, + num_examples=healthbench_num_examples, n_repeats=1, n_threads=args.n_threads or 1, subset_name="consensus", ) case "aime25": return AIME25Eval( - n_repeats=1 if args.debug else 8, + n_repeats=n_repeats, num_examples=num_examples, n_threads=args.n_threads or 1, ) diff --git a/gpt_oss/evals/cli_utils.py b/gpt_oss/evals/cli_utils.py new file mode 100644 index 00000000..911b8fa0 --- /dev/null +++ b/gpt_oss/evals/cli_utils.py @@ -0,0 +1,17 @@ +def resolve_num_examples( + explicit_examples: int | None, debug_mode: bool, debug_default: int +) -> int | None: + if explicit_examples is not None: + if explicit_examples == 0: + return debug_default if debug_mode else None + return explicit_examples + return debug_default if debug_mode else None + + +def resolve_n_repeats(explicit_examples: int | None, debug_mode: bool) -> int: + explicit_subset = explicit_examples is not None and explicit_examples != 0 + return 1 if explicit_subset or debug_mode else 8 + + +def resolve_gpqa_debug_mode(explicit_examples: int | None, debug_mode: bool) -> bool: + return debug_mode and explicit_examples is None diff --git a/tests/gpt_oss/evals/test_cli_examples.py b/tests/gpt_oss/evals/test_cli_examples.py new file mode 100644 index 00000000..f141d2c8 --- /dev/null +++ b/tests/gpt_oss/evals/test_cli_examples.py @@ -0,0 +1,47 @@ +from gpt_oss.evals.cli_utils import ( + resolve_gpqa_debug_mode, + resolve_n_repeats, + resolve_num_examples, +) + + +def test_explicit_examples_override_debug_default() -> None: + assert resolve_num_examples(2, debug_mode=True, debug_default=10) == 2 + + +def test_debug_default_is_used_without_explicit_examples() -> None: + assert resolve_num_examples(None, debug_mode=True, debug_default=10) == 10 + + +def test_full_eval_keeps_unlimited_examples() -> None: + assert resolve_num_examples(None, debug_mode=False, debug_default=10) is None + + +def test_zero_examples_keeps_full_run_outside_debug() -> None: + assert resolve_num_examples(0, debug_mode=False, debug_default=10) is None + + +def test_zero_examples_uses_debug_default_in_debug_mode() -> None: + assert resolve_num_examples(0, debug_mode=True, debug_default=10) == 10 + + +def test_explicit_subset_uses_single_repeat() -> None: + assert resolve_n_repeats(2, debug_mode=False) == 1 + + +def test_debug_run_uses_single_repeat() -> None: + assert resolve_n_repeats(None, debug_mode=True) == 1 + + +def test_full_run_keeps_eight_repeats() -> None: + assert resolve_n_repeats(None, debug_mode=False) == 8 + + +def test_zero_examples_keeps_full_run_repeats() -> None: + assert resolve_n_repeats(0, debug_mode=False) == 8 + + +def test_gpqa_fixed_debug_example_is_default_only() -> None: + assert resolve_gpqa_debug_mode(None, debug_mode=True) is True + assert resolve_gpqa_debug_mode(2, debug_mode=True) is False + assert resolve_gpqa_debug_mode(None, debug_mode=False) is False