369 lines
11 KiB
Python
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}
|