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
153 changes: 153 additions & 0 deletions src/evaluation/ab.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,153 @@
"""Compare two configurations without the three mistakes that keep recurring.

Every latency and retrieval comparison in this repository has gone wrong in
one of three ways, more than once each, and knowing the lesson has not
prevented repeating it within hours. So the checks live here rather than in
whoever is writing the script.

**1. Arms measured in sequence.** The second arm benefits from a warm process,
a warm connection and whatever the API is doing that minute. Twice this has
produced a result in the direction the author was hoping for -- and once, in
an earlier measurement, showed a real improvement as a 2.5s regression. `run`
alternates the arms and never offers a mode that does not.

**2. The manipulation never verified.** A prompt asking for "exactly 1"
alternate produced four; a ContextVar set around a call was overwritten by the
node inside it. Both looked like ordinary results. Each arm must declare a
`precondition`, and `run` refuses to report a comparison whose precondition
did not hold rather than printing numbers that mean nothing.

**3. A precondition that cannot fail.** The check above is only worth the
discrimination in it: `lambda: True` passes every time and reads like a guard.
So after each sample the *other* arms' preconditions are evaluated against the
configuration that is actually active, and if one of them also holds, the two
preconditions do not tell the arms apart and the comparison is refused. This
is the same rule the rest of the repository learned the hard way -- an
assertion that something is absent proves nothing until the same check has
shown it can be present.

**4. A percentile the sample cannot support.** p90 from fifteen samples is
about the second-highest value. `Summary` reports p50 with min and max, and
says so, instead of implying a tail estimate that is not there.
"""

import statistics
from collections.abc import Awaitable, Callable, Iterable
from dataclasses import dataclass, field
from typing import Generic, TypeVar

T = TypeVar("T")

#: Below this, a median is a gesture. Not enforced -- a small run is often the
#: right thing -- but reported, so nobody quotes six samples as a p50.
ADVISORY_MIN_SAMPLES = 10


@dataclass
class Arm(Generic[T]):
"""One configuration under test.

`apply` makes the configuration active and is called immediately before
each sample, never once per arm -- a setting applied once and mutated by
something else in between is exactly the failure this guards.

`precondition` is checked after each sample and must return True for the
manipulation to be believed. Returning True unconditionally defeats the
purpose, so write it against something the configuration actually changes.
"""

name: str
apply: Callable[[], None]
precondition: Callable[[], bool]
samples: list[float] = field(default_factory=list)
precondition_failures: int = 0
#: Times another arm's precondition also held under this arm's
#: configuration, meaning the two do not discriminate.
indiscriminate: int = 0

@property
def held(self) -> bool:
return (
self.precondition_failures == 0
and self.indiscriminate == 0
and bool(self.samples)
)


@dataclass
class Summary:
arms: list[Arm[float]]

@property
def trustworthy(self) -> bool:
return all(arm.held for arm in self.arms)

def report(self) -> str:
lines = []
for arm in self.arms:
if not arm.samples:
lines.append(f" {arm.name:24s} no samples")
continue
values = sorted(arm.samples)
note = ""
if arm.precondition_failures:
note = (
f" PRECONDITION FAILED on {arm.precondition_failures} "
f"of {len(values)} samples -- this comparison means nothing"
)
elif arm.indiscriminate:
note = (
f" PRECONDITION DOES NOT DISCRIMINATE: another arm's held "
f"under this one's configuration on {arm.indiscriminate} "
f"samples, so it cannot tell the arms apart"
)
elif len(values) < ADVISORY_MIN_SAMPLES:
note = f" (n={len(values)}, small; treat the median loosely)"
lines.append(
f" {arm.name:24s} p50 {statistics.median(values):6.2f} "
f"min {values[0]:6.2f} max {values[-1]:6.2f} n={len(values)}{note}"
)
if not self.trustworthy:
lines.append(
" Refusing to draw a conclusion: an arm's precondition did not hold."
)
return "\n".join(lines)


async def run(
arms: list[Arm[float]],
cases: Iterable[object],
measure: Callable[[object, str], Awaitable[float]],
*,
repeats: int = 5,
) -> Summary:
"""Measure every case under every arm, alternating which arm goes first.

The alternation is not optional and the order is derived from the case and
repeat indices, so a run is deterministic in structure while still
splitting any warming effect evenly between the arms.
"""
if len(arms) < 2:
raise ValueError("an A/B comparison needs at least two arms")
if any(arm.samples for arm in arms):
# Re-using arms silently mixes two runs into one distribution, and the
# result looks like an ordinary noisy measurement.
raise ValueError("these arms already hold samples; build fresh ones")
cases = list(cases)
for repeat in range(repeats):
for index, case in enumerate(cases):
# Rotate rather than reverse: with more than two arms, reversing
# leaves the middle one always in the middle.
offset = (repeat + index) % len(arms)
ordered = arms[offset:] + arms[:offset]
for arm in ordered:
arm.apply()
arm.samples.append(await measure(case, arm.name))
if not arm.precondition():
arm.precondition_failures += 1
continue
# The configuration for `arm` is still active, so any other
# arm whose precondition also holds is not distinguishing.
if any(other is not arm and other.precondition() for other in arms):
arm.indiscriminate += 1
return Summary(arms)
168 changes: 168 additions & 0 deletions tests/evaluation/test_ab.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,168 @@
"""The harness exists to catch three mistakes, so the tests are that it does.

Each test recreates a mistake actually made in this repository rather than a
hypothetical one.
"""

import asyncio

import pytest

from evaluation.ab import ADVISORY_MIN_SAMPLES, Arm, run

#: Which arm's configuration is currently applied. Shared, because a
#: precondition that does not read something the arms actually change cannot
#: tell them apart -- which the harness now refuses, and which this helper
#: used to do by returning a constant.
ACTIVE: list[str] = []


def _arm(name: str, holds: bool = True) -> tuple[Arm[float], list[str]]:
applied: list[str] = []

def apply() -> None:
applied.append(name)
ACTIVE.clear()
ACTIVE.append(name)

return (
Arm(
name=name,
apply=apply,
precondition=(lambda: [name] == ACTIVE) if holds else (lambda: False),
),
applied,
)


def _always_true_arm(name: str) -> Arm[float]:
"""A guard that cannot fail, which is the thing the harness must catch."""
return Arm(name=name, apply=lambda: None, precondition=lambda: True)


def _measure_order(order: list[str]): # type: ignore[no-untyped-def]
async def measure(_case: object, arm: str) -> float:
order.append(arm)
return 1.0

return measure


def test_the_arms_alternate_rather_than_running_in_sequence() -> None:
# The mistake: measuring arm A fully, then arm B, so B inherits a warm
# process and whatever the API is doing. It has twice produced a result
# in the direction the author wanted.
a, _ = _arm("a")
b, _ = _arm("b")
order: list[str] = []
asyncio.run(run([a, b], ["q1", "q2"], _measure_order(order), repeats=2))
first_of_each_pair = order[::2]
assert set(first_of_each_pair) == {"a", "b"}, "one arm always went first"


def test_a_failed_precondition_refuses_the_comparison() -> None:
# The mistake: a prompt asked for "exactly 1" alternate and got four; a
# ContextVar was overwritten by the node it was set around. Both produced
# ordinary-looking numbers that measured the baseline twice.
good, _ = _arm("good")
broken, _ = _arm("broken", holds=False)
summary = asyncio.run(run([good, broken], ["q"], _measure_order([]), repeats=3))
assert not summary.trustworthy
assert "PRECONDITION FAILED" in summary.report()
assert "means nothing" in summary.report()
assert "Refusing to draw a conclusion" in summary.report()


def test_a_held_precondition_reports_normally() -> None:
a, _ = _arm("a")
b, _ = _arm("b")
summary = asyncio.run(run([a, b], ["q"], _measure_order([]), repeats=6))
assert summary.trustworthy
assert "PRECONDITION FAILED" not in summary.report()
assert "Refusing" not in summary.report()


def test_the_configuration_is_applied_per_sample_not_once_per_arm() -> None:
# A setting applied once and mutated in between is the same class of
# failure as never applying it: the numbers look fine either way.
a, applied_a = _arm("a")
b, _ = _arm("b")
asyncio.run(run([a, b], ["q1", "q2"], _measure_order([]), repeats=3))
assert len(applied_a) == len(a.samples) == 6


def test_a_small_sample_is_flagged_rather_than_quoted_confidently() -> None:
# p90 from fifteen samples is about the second-highest value, and one was
# quoted from fifteen. The harness reports min and max instead, and says
# when the median is thin.
a, _ = _arm("a")
b, _ = _arm("b")
summary = asyncio.run(run([a, b], ["q"], _measure_order([]), repeats=2))
assert all(len(arm.samples) < ADVISORY_MIN_SAMPLES for arm in summary.arms)
assert "small; treat the median loosely" in summary.report()
assert "p90" not in summary.report()


def test_one_arm_is_not_a_comparison() -> None:
a, _ = _arm("a")
with pytest.raises(ValueError, match="at least two arms"):
asyncio.run(run([a], ["q"], _measure_order([]), repeats=1))


def test_a_precondition_that_cannot_fail_is_caught() -> None:
# The hole this harness had: `lambda: True` reads like a guard, passes
# every sample, and proves nothing. Same shape as an absence test that
# never showed the thing can be present -- which is the mistake this file
# was written to stop, present in the file itself.
summary = asyncio.run(
run(
[_always_true_arm("always-a"), _always_true_arm("always-b")],
["q"],
_measure_order([]),
repeats=3,
)
)
assert not summary.trustworthy
assert "DOES NOT DISCRIMINATE" in summary.report()
assert "cannot tell the arms apart" in summary.report()


def test_discriminating_preconditions_are_accepted() -> None:
# The honest case: a shared flag the arms actually set, so each arm's
# precondition is false under the other's configuration.
active: list[str] = []

def arm(name: str) -> Arm[float]:
return Arm(
name=name,
apply=lambda: active.clear() or active.append(name), # type: ignore[func-returns-value]
precondition=lambda: active == [name],
)

summary = asyncio.run(
run([arm("a"), arm("b")], ["q"], _measure_order([]), repeats=6)
)
assert summary.trustworthy
assert "DISCRIMINATE" not in summary.report()


def test_reusing_arms_across_runs_is_refused() -> None:
# Two runs into one distribution looks exactly like one noisy run.
a, _ = _arm("a")
b, _ = _arm("b")
asyncio.run(run([a, b], ["q"], _measure_order([]), repeats=2))
with pytest.raises(ValueError, match="already hold samples"):
asyncio.run(run([a, b], ["q"], _measure_order([]), repeats=2))


def test_three_arms_each_take_every_position() -> None:
# Reversing gives two orderings, so with three arms the middle one is
# always in the middle and carries a systematic warming bias.
order: list[str] = []
arms = [_arm(name)[0] for name in ("a", "b", "c")]
asyncio.run(run(arms, ["q"], _measure_order(order), repeats=3))
positions: dict[str, set[int]] = {name: set() for name in ("a", "b", "c")}
for start in range(0, len(order), 3):
for position, name in enumerate(order[start : start + 3]):
positions[name].add(position)
assert all(len(seen) == 3 for seen in positions.values()), positions
Loading