Files
6krrt/tests/test_exploration.py

369 lines
11 KiB
Python

"""Tests for exploration.py — epsilon-greedy exploration chooser.
Every test follows Given/When/Then. One ``When`` per test; if more than one
condition needs coverage, split into separate tests.
"""
from __future__ import annotations
import random
import pytest
from exploration import choose
# --- helpers -------------------------------------------------------------- #
WINNER = {"model_id": "a1", "cost": 10.0, "quality": 0.95}
ALTERNATIVE_FEW = {"model_id": "b1", "cost": 8.0, "quality": 0.90, "outcome_samples": 2}
ALTERNATIVE_MANY = {"model_id": "c1", "cost": 9.0, "quality": 0.92, "outcome_samples": 50}
EXPLORATION_KW = {"epsilon": 0.5, "max_cost_ratio": 1.5, "rng": None}
class FakeRandom:
"""Deterministic random that always returns one fixed value."""
def __init__(self, value: float) -> None:
self.value = value
def random(self) -> float:
return self.value
def _choose(
ranked: list[dict],
sample_counts: dict[str, int],
**over,
) -> tuple[dict, bool]:
kw = {**EXPLORATION_KW, **over}
r = kw.pop("rng") or FakeRandom(1.0)
return choose(
ranked,
sample_counts,
epsilon=kw["epsilon"],
max_cost_ratio=kw["max_cost_ratio"],
rng=r,
)
# --- fallthrough: epsilon <= 0 -------------------------------------------- #
def test_zero_epsilon_returns_winner_no_exploration():
result, explored = _choose([WINNER], {"a1": 100})
assert result["model_id"] == "a1"
assert explored is False
def test_negative_epsilon_also_disabled():
result, explored = _choose(
[WINNER, ALTERNATIVE_FEW], {"a1": 100, "b1": 2}, epsilon=-0.1
)
assert result["model_id"] == "a1"
assert explored is False
# --- fallthrough: singleton eligible set ---------------------------------- #
def test_singleton_returns_winner():
result, explored = _choose([{"model_id": "x"}], {"x": 5})
assert result["model_id"] == "x"
assert explored is False
# --- fallthrough: exploitation (rng >= epsilon) --------------------------- #
def test_rng_above_epsilon_returns_winner():
# Given: rng.random() returns 0.7 > epsilon=0.5.
# When: 0.7 >= 0.5 -> exploitation.
# Then: winner returned.
result, explored = _choose(
[WINNER, ALTERNATIVE_FEW],
{"a1": 100, "b1": 2},
rng=FakeRandom(0.7),
)
assert result["model_id"] == "a1"
assert explored is False
def test_rng_at_exactly_epsilon_returns_winner():
# Given: rng.random() returns exactly epsilon.
# When: 0.5 >= 0.5 -> exploitation.
result, explored = _choose(
[WINNER, ALTERNATIVE_FEW],
{"a1": 100, "b1": 2},
epsilon=0.5,
rng=FakeRandom(0.5),
)
assert result["model_id"] == "a1"
assert explored is False
# --- exploration (rng < epsilon) ------------------------------------------ #
def test_rng_just_below_epsilon_triggers_exploration():
# Given: rng.random() returns 0.49 < epsilon=0.5.
# When: exploration path is entered.
# Then: the least-evidenced alternative is chosen.
result, explored = _choose(
[WINNER, ALTERNATIVE_FEW],
{"a1": 100, "b1": 2},
epsilon=0.5,
rng=FakeRandom(0.49),
)
assert result["model_id"] == "b1"
assert explored is True
# --- exploration: selects fewest-sampled alternative ---------------------- #
def test_fewest_samples_wins_exploration():
ranked = [WINNER, ALTERNATIVE_FEW, ALTERNATIVE_MANY]
result, explored = _choose(
ranked, {"a1": 100, "b1": 2, "c1": 50}, rng=FakeRandom(0.1)
)
assert result["model_id"] == "b1"
assert explored is True
def test_all_alternatives_have_zero_samples_uses_cost_tiebreak():
ranked = [
{"model_id": "a1", "cost": 10.0},
{"model_id": "b1", "cost": 8.0},
{"model_id": "c1", "cost": 6.0},
]
result, explored = _choose(
ranked, {"a1": 0, "b1": 0, "c1": 0},
epsilon=0.5, rng=FakeRandom(0.1), max_cost_ratio=2.0
)
assert result["model_id"] == "c1" # cheapest gets picked
assert explored is True
def test_fewest_samples_beats_lower_cost():
ranked = [
{"model_id": "a1", "cost": 10.0},
{"model_id": "b1", "cost": 5.0}, # cheaper
{"model_id": "c1", "cost": 7.0}, # more expensive but fewer samples
]
result, explored = _choose(
ranked,
{"a1": 100, "b1": 50, "c1": 1},
epsilon=0.5,
rng=FakeRandom(0.1),
max_cost_ratio=2.0,
)
assert result["model_id"] == "c1"
assert explored is True
# --- cost ratio constraint ----------------------------------------------- #
def test_alt_exceeds_cost_ratio_returns_winner():
alt = {"model_id": "b1", "cost": 30.0}
result, explored = _choose(
[WINNER, alt],
{"a1": 100, "b1": 1},
rng=FakeRandom(0.1),
max_cost_ratio=1.5,
)
assert result["model_id"] == "a1"
assert explored is False
def test_alt_exactly_at_cost_ratio_is_allowed():
alt = {"model_id": "b1", "cost": 15.0}
result, explored = _choose(
[WINNER, alt],
{"a1": 100, "b1": 1},
rng=FakeRandom(0.1),
max_cost_ratio=1.5,
)
assert result["model_id"] == "b1"
assert explored is True
def test_alt_just_under_cost_ratio_keeps_alternative():
# Given: alternative cost is just under the ratio bound.
# When: 14.99 <= 1.5 * 10.0.
# Then: alternative is selected.
alt = {"model_id": "b1", "cost": 14.99}
result, explored = _choose(
[WINNER, alt],
{"a1": 100, "b1": 1},
rng=FakeRandom(0.1),
max_cost_ratio=1.5,
)
assert result["model_id"] == "b1"
assert explored is True
def test_multiple_alts_some_filtered_by_ratio_picks_remaining_best():
alt_b = {"model_id": "b1", "cost": 12.0} # within 1.5x
alt_c = {"model_id": "c1", "cost": 50.0} # way over
ranked = [WINNER, alt_b, alt_c]
result, explored = _choose(
ranked,
{"a1": 100, "b1": 0, "c1": 0},
rng=FakeRandom(0.1),
max_cost_ratio=1.5,
)
assert result["model_id"] == "b1"
assert explored is True
# None cost handling ------------------------------------------------------- #
def test_winner_none_cost_no_ratio_check():
winner_none = {"model_id": "a1", "cost": None, "quality": 0.90}
alt = {"model_id": "b1", "cost": 100.0}
result, explored = _choose(
[winner_none, alt],
{"a1": 100, "b1": 0},
rng=FakeRandom(0.1),
max_cost_ratio=1.0, # would reject if ratio check ran
)
assert result["model_id"] == "b1"
assert explored is True
def test_alternative_none_cost_skips_ratio_check():
winner = {"model_id": "a1", "cost": 10.0}
alt = {"model_id": "b1", "cost": None}
result, explored = _choose(
[winner, alt],
{"a1": 100, "b1": 0},
rng=FakeRandom(0.1),
max_cost_ratio=1.0,
)
assert result["model_id"] == "b1"
assert explored is True
def test_both_none_costs_allows_exploration():
ranked = [
{"model_id": "a1", "cost": None},
{"model_id": "b1", "cost": None},
]
result, explored = _choose(
ranked, {"a1": 10, "b1": 0},
rng=FakeRandom(0.1),
max_cost_ratio=1.0,
)
assert result["model_id"] == "b1"
assert explored is True
# --- sample_counts edge cases --------------------------------------------- #
def test_alternative_absent_from_sample_counts_gets_zero():
result, explored = _choose(
[WINNER, ALTERNATIVE_FEW],
{"a1": 100}, # b1 missing from sample_counts
rng=FakeRandom(0.1),
max_cost_ratio=2.0,
)
assert result["model_id"] == "b1"
assert explored is True
def test_empty_sample_counts_dict_works():
ranked = [
{"model_id": "a1", "cost": 10.0},
{"model_id": "b1", "cost": 12.0},
{"model_id": "c1", "cost": 6.0},
]
result, explored = _choose(
ranked, {}, # empty sample_counts
rng=FakeRandom(0.1),
max_cost_ratio=2.0,
)
assert result["model_id"] == "c1"
assert explored is True
# --- pure function integrity ---------------------------------------------- #
def test_deterministic_with_same_seed():
ranked = [WINNER, ALTERNATIVE_FEW]
rng = FakeRandom(0.1)
r1 = choose(ranked, {"a1": 100, "b1": 1}, epsilon=0.5, max_cost_ratio=2.0, rng=rng)
rng2 = FakeRandom(0.1)
r2 = choose(ranked, {"a1": 100, "b1": 1}, epsilon=0.5, max_cost_ratio=2.0, rng=rng2)
assert r1 == r2
def test_returns_tuple_of_two_elements():
for rng_val in (0.0, 0.49, 0.5, 0.99):
rng = FakeRandom(rng_val)
result = choose(
[WINNER, ALTERNATIVE_FEW],
{"a1": 100, "b1": 1},
epsilon=0.5,
max_cost_ratio=2.0,
rng=rng,
)
assert isinstance(result, tuple)
assert len(result) == 2
assert isinstance(result[0], dict)
assert isinstance(result[1], bool)
# --- property invariants -------------------------------------------------- #
def test_returned_row_is_always_an_element_of_ranked():
for rng_val in (0.0, 0.1, 0.5, 0.9, 0.99):
rng = FakeRandom(rng_val)
ranked = [WINNER, ALTERNATIVE_FEW, ALTERNATIVE_MANY]
row, _ = _choose(ranked, {"a1": 100, "b1": 2, "c1": 50}, rng=rng)
assert row in ranked
def test_exploration_cost_does_not_exceed_ratio_cap():
winner = {"model_id": "a", "cost": 10.0}
within = {"model_id": "b", "cost": 14.0}
over = {"model_id": "c", "cost": 25.0}
rng = random.Random(42)
for _ in range(500):
row, was_exploration = choose(
[winner, within, over],
{"a": 0, "b": 0, "c": 0},
epsilon=0.5,
max_cost_ratio=1.5,
rng=rng,
)
assert row in (winner, within, over)
if was_exploration:
assert row["cost"] <= 1.5 * winner["cost"]
def test_seeded_rng_explore_share_lands_within_tolerance_of_epsilon():
rng = random.Random(12345)
ranked = [WINNER, ALTERNATIVE_FEW, ALTERNATIVE_MANY]
sample_counts = {"a1": 100, "b1": 2, "c1": 50}
epsilon = 0.25
trials = 4000
explored = sum(
choose(ranked, sample_counts, epsilon=epsilon, max_cost_ratio=2.0, rng=rng)[1]
for _ in range(trials)
)
assert explored / trials == pytest.approx(epsilon, abs=0.03)
def test_exploration_on_ranked_two_excludes_winner_and_first_runner_up():
ranked = [WINNER, ALTERNATIVE_FEW, ALTERNATIVE_MANY]
# Make ranked[2] the least-evidenced alternative.
sample_counts = {"a1": 100, "b1": 50, "c1": 1}
result, was_exploration = choose(
ranked,
sample_counts,
epsilon=1.0,
max_cost_ratio=2.0,
rng=FakeRandom(0.0),
)
assert was_exploration is True
assert result is ranked[2]
winner_id = ranked[0]["model_id"]
runner_up_id = ranked[1]["model_id"]
assert result["model_id"] not in {winner_id, runner_up_id}