diff --git a/src/evalwire/__init__.py b/src/evalwire/__init__.py index 0c06b8d..25c6ce3 100644 --- a/src/evalwire/__init__.py +++ b/src/evalwire/__init__.py @@ -4,6 +4,8 @@ from importlib.metadata import version from evalwire.evaluators import ( + make_all_pass_evaluator, + make_any_pass_evaluator, make_contains_evaluator, make_exact_match_evaluator, make_json_match_evaluator, @@ -13,9 +15,9 @@ make_regex_evaluator, make_schema_evaluator, make_top_k_evaluator, + make_weighted_evaluator, ) from evalwire.observability import setup_observability -from evalwire.results import ResultCollector from evalwire.runner import ExperimentRunner from evalwire.uploader import DatasetUploader @@ -26,7 +28,8 @@ __all__ = [ "DatasetUploader", "ExperimentRunner", - "ResultCollector", + "make_all_pass_evaluator", + "make_any_pass_evaluator", "make_contains_evaluator", "make_exact_match_evaluator", "make_json_match_evaluator", @@ -36,10 +39,8 @@ "make_regex_evaluator", "make_schema_evaluator", "make_top_k_evaluator", + "make_weighted_evaluator", "setup_observability", - # LangGraph helpers — available when the `evalwire[langgraph]` extra is - # installed. Importing from the top-level package is supported; if - # langgraph is absent the ImportError is raised at call time, not import time. "build_subgraph", "invoke_node", ] diff --git a/src/evalwire/evaluators/__init__.py b/src/evalwire/evaluators/__init__.py index f087638..7cff760 100644 --- a/src/evalwire/evaluators/__init__.py +++ b/src/evalwire/evaluators/__init__.py @@ -1,5 +1,10 @@ """Built-in evaluator factories for evalwire.""" +from evalwire.evaluators.composition import ( + make_all_pass_evaluator, + make_any_pass_evaluator, + make_weighted_evaluator, +) from evalwire.evaluators.contains import make_contains_evaluator from evalwire.evaluators.exact_match import make_exact_match_evaluator from evalwire.evaluators.json_match import make_json_match_evaluator @@ -11,6 +16,8 @@ from evalwire.evaluators.top_k import make_top_k_evaluator __all__ = [ + "make_all_pass_evaluator", + "make_any_pass_evaluator", "make_contains_evaluator", "make_exact_match_evaluator", "make_json_match_evaluator", @@ -20,4 +27,5 @@ "make_regex_evaluator", "make_schema_evaluator", "make_top_k_evaluator", + "make_weighted_evaluator", ] diff --git a/src/evalwire/evaluators/composition.py b/src/evalwire/evaluators/composition.py new file mode 100644 index 0000000..e0ce6bd --- /dev/null +++ b/src/evalwire/evaluators/composition.py @@ -0,0 +1,93 @@ +"""Evaluator composition factories.""" + +from collections.abc import Callable + + +def make_weighted_evaluator( + evaluators: list[tuple[Callable, float]], +) -> Callable[[str, dict], float]: + """Return a weighted-average composition evaluator. + + Parameters + ---------- + evaluators: + List of ``(evaluator, weight)`` pairs. Weights are normalised + internally and must be non-negative with at least one positive value. + + Returns + ------- + Callable[[str, dict], float] + Evaluator with signature ``weighted(output, expected) -> float``. + """ + if not evaluators: + raise ValueError("at least one evaluator is required") + for _, w in evaluators: + if w < 0: + raise ValueError("weights must be non-negative") + total = sum(w for _, w in evaluators) + if total == 0: + raise ValueError("total weight must be non-zero") + + normalised = [(fn, w / total) for fn, w in evaluators] + + def weighted(output: str, expected: dict) -> float: + return float(sum(fn(output, expected) * w for fn, w in normalised)) + + weighted.__name__ = "weighted" + return weighted + + +def make_all_pass_evaluator( + evaluators: list[Callable], +) -> Callable[[str, dict], bool]: + """Return an AND-composition evaluator. + + Returns ``True`` only if every sub-evaluator returns a truthy value. + Short-circuits on the first falsy result. + + Parameters + ---------- + evaluators: + Non-empty list of evaluator callables. + + Returns + ------- + Callable[[str, dict], bool] + Evaluator with signature ``all_pass(output, expected) -> bool``. + """ + if not evaluators: + raise ValueError("at least one evaluator is required") + + def all_pass(output: str, expected: dict) -> bool: + return all(bool(fn(output, expected)) for fn in evaluators) + + all_pass.__name__ = "all_pass" + return all_pass + + +def make_any_pass_evaluator( + evaluators: list[Callable], +) -> Callable[[str, dict], bool]: + """Return an OR-composition evaluator. + + Returns ``True`` if at least one sub-evaluator returns a truthy value. + Short-circuits on the first truthy result. + + Parameters + ---------- + evaluators: + Non-empty list of evaluator callables. + + Returns + ------- + Callable[[str, dict], bool] + Evaluator with signature ``any_pass(output, expected) -> bool``. + """ + if not evaluators: + raise ValueError("at least one evaluator is required") + + def any_pass(output: str, expected: dict) -> bool: + return any(bool(fn(output, expected)) for fn in evaluators) + + any_pass.__name__ = "any_pass" + return any_pass diff --git a/tests/test_composition_evaluators.py b/tests/test_composition_evaluators.py new file mode 100644 index 0000000..12bcdcc --- /dev/null +++ b/tests/test_composition_evaluators.py @@ -0,0 +1,185 @@ +"""Tests for evalwire.evaluators.composition factories.""" + +import pytest + +from evalwire.evaluators.composition import ( + make_all_pass_evaluator, + make_any_pass_evaluator, + make_weighted_evaluator, +) + + +def _always(value): + def evaluator(output, expected): + return value + + return evaluator + + +TRUE_EVAL = _always(True) +FALSE_EVAL = _always(False) +ONE_EVAL = _always(1.0) +ZERO_EVAL = _always(0.0) +HALF_EVAL = _always(0.5) + +EXPECTED = {"expected_output": ["x"]} + + +class TestMakeWeightedEvaluator: + def test_function_name(self): + fn = make_weighted_evaluator([(ONE_EVAL, 1.0)]) + assert fn.__name__ == "weighted" # ty: ignore[unresolved-attribute] + + def test_single_evaluator_returns_its_score(self): + fn = make_weighted_evaluator([(ONE_EVAL, 1.0)]) + assert fn("out", EXPECTED) == pytest.approx(1.0) + + def test_weights_are_normalised(self): + fn = make_weighted_evaluator([(ONE_EVAL, 2.0), (ZERO_EVAL, 2.0)]) + assert fn("out", EXPECTED) == pytest.approx(0.5) + + def test_equal_weights_averages_scores(self): + fn = make_weighted_evaluator([(ONE_EVAL, 1.0), (ZERO_EVAL, 1.0)]) + assert fn("out", EXPECTED) == pytest.approx(0.5) + + def test_bool_scores_treated_as_float(self): + fn = make_weighted_evaluator([(TRUE_EVAL, 1.0), (FALSE_EVAL, 1.0)]) + assert fn("out", EXPECTED) == pytest.approx(0.5) + + def test_unequal_weights_produce_correct_result(self): + fn = make_weighted_evaluator([(ONE_EVAL, 0.7), (ZERO_EVAL, 0.3)]) + assert fn("out", EXPECTED) == pytest.approx(0.7) + + def test_return_type_is_float(self): + fn = make_weighted_evaluator([(TRUE_EVAL, 1.0)]) + result = fn("out", EXPECTED) + assert isinstance(result, float) + + def test_empty_list_raises_value_error(self): + with pytest.raises(ValueError, match="at least one"): + make_weighted_evaluator([]) + + def test_negative_weight_raises_value_error(self): + with pytest.raises(ValueError, match="non-negative"): + make_weighted_evaluator([(ONE_EVAL, -1.0)]) + + def test_all_zero_weights_raises_value_error(self): + with pytest.raises(ValueError, match="non-zero"): + make_weighted_evaluator([(ONE_EVAL, 0.0), (ZERO_EVAL, 0.0)]) + + def test_passes_output_and_expected_to_sub_evaluators(self): + received = [] + + def recording_eval(output, expected): + received.append((output, expected)) + return 1.0 + + fn = make_weighted_evaluator([(recording_eval, 1.0)]) + fn("hello", EXPECTED) + assert received == [("hello", EXPECTED)] + + +class TestMakeAllPassEvaluator: + def test_function_name(self): + fn = make_all_pass_evaluator([TRUE_EVAL]) + assert fn.__name__ == "all_pass" # ty: ignore[unresolved-attribute] + + def test_all_true_returns_true(self): + fn = make_all_pass_evaluator([TRUE_EVAL, TRUE_EVAL]) + assert fn("out", EXPECTED) is True + + def test_one_false_returns_false(self): + fn = make_all_pass_evaluator([TRUE_EVAL, FALSE_EVAL]) + assert fn("out", EXPECTED) is False + + def test_all_false_returns_false(self): + fn = make_all_pass_evaluator([FALSE_EVAL, FALSE_EVAL]) + assert fn("out", EXPECTED) is False + + def test_single_true_evaluator(self): + fn = make_all_pass_evaluator([TRUE_EVAL]) + assert fn("out", EXPECTED) is True + + def test_single_false_evaluator(self): + fn = make_all_pass_evaluator([FALSE_EVAL]) + assert fn("out", EXPECTED) is False + + def test_truthy_float_counts_as_pass(self): + fn = make_all_pass_evaluator([ONE_EVAL]) + assert fn("out", EXPECTED) is True + + def test_zero_float_counts_as_fail(self): + fn = make_all_pass_evaluator([ZERO_EVAL]) + assert fn("out", EXPECTED) is False + + def test_return_type_is_bool(self): + fn = make_all_pass_evaluator([TRUE_EVAL]) + assert isinstance(fn("out", EXPECTED), bool) + + def test_empty_list_raises_value_error(self): + with pytest.raises(ValueError, match="at least one"): + make_all_pass_evaluator([]) + + def test_short_circuits_on_first_failure(self): + called = [] + + def recording_eval(output, expected): + called.append(1) + return True + + fn = make_all_pass_evaluator([FALSE_EVAL, recording_eval]) + fn("out", EXPECTED) + assert called == [] + + +class TestMakeAnyPassEvaluator: + def test_function_name(self): + fn = make_any_pass_evaluator([TRUE_EVAL]) + assert fn.__name__ == "any_pass" # ty: ignore[unresolved-attribute] + + def test_all_true_returns_true(self): + fn = make_any_pass_evaluator([TRUE_EVAL, TRUE_EVAL]) + assert fn("out", EXPECTED) is True + + def test_one_true_returns_true(self): + fn = make_any_pass_evaluator([FALSE_EVAL, TRUE_EVAL]) + assert fn("out", EXPECTED) is True + + def test_all_false_returns_false(self): + fn = make_any_pass_evaluator([FALSE_EVAL, FALSE_EVAL]) + assert fn("out", EXPECTED) is False + + def test_single_true_evaluator(self): + fn = make_any_pass_evaluator([TRUE_EVAL]) + assert fn("out", EXPECTED) is True + + def test_single_false_evaluator(self): + fn = make_any_pass_evaluator([FALSE_EVAL]) + assert fn("out", EXPECTED) is False + + def test_truthy_float_counts_as_pass(self): + fn = make_any_pass_evaluator([HALF_EVAL]) + assert fn("out", EXPECTED) is True + + def test_zero_float_counts_as_fail(self): + fn = make_any_pass_evaluator([ZERO_EVAL]) + assert fn("out", EXPECTED) is False + + def test_return_type_is_bool(self): + fn = make_any_pass_evaluator([FALSE_EVAL]) + assert isinstance(fn("out", EXPECTED), bool) + + def test_empty_list_raises_value_error(self): + with pytest.raises(ValueError, match="at least one"): + make_any_pass_evaluator([]) + + def test_short_circuits_on_first_success(self): + called = [] + + def recording_eval(output, expected): + called.append(1) + return False + + fn = make_any_pass_evaluator([TRUE_EVAL, recording_eval]) + fn("out", EXPECTED) + assert called == []