Skip to content
Merged
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
11 changes: 6 additions & 5 deletions src/evalwire/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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

Expand All @@ -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",
Expand All @@ -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",
]
Expand Down
8 changes: 8 additions & 0 deletions src/evalwire/evaluators/__init__.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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",
Expand All @@ -20,4 +27,5 @@
"make_regex_evaluator",
"make_schema_evaluator",
"make_top_k_evaluator",
"make_weighted_evaluator",
]
93 changes: 93 additions & 0 deletions src/evalwire/evaluators/composition.py
Original file line number Diff line number Diff line change
@@ -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
185 changes: 185 additions & 0 deletions tests/test_composition_evaluators.py
Original file line number Diff line number Diff line change
@@ -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 == []