Skip to content

Commit 43372e6

Browse files
authored
Merge pull request #28 from scope-lab-vu/bug/update_function_tests
more tests for update fns and schedulers
2 parents 545eb32 + 6f7406e commit 43372e6

2 files changed

Lines changed: 217 additions & 48 deletions

File tree

tests/test_schedulers.py

Lines changed: 118 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -588,5 +588,123 @@ def test_many_windows(self):
588588
assert sched(14) is False
589589

590590

591+
# ============================================================
592+
# Parametrized common-property tests for all Schedulers
593+
# ============================================================
594+
595+
_scheduler_factories = [
596+
pytest.param(
597+
lambda: ContinuousScheduler(start=10, end=50),
598+
id="ContinuousScheduler",
599+
),
600+
pytest.param(
601+
lambda: DiscreteScheduler(event_list={15, 30, 45}, start=10, end=50),
602+
id="DiscreteScheduler",
603+
),
604+
pytest.param(
605+
lambda: PeriodicScheduler(period=5, start=10, end=50),
606+
id="PeriodicScheduler",
607+
),
608+
pytest.param(
609+
lambda: RandomScheduler(probability=0.5, start=10, end=50, seed=42),
610+
id="RandomScheduler",
611+
),
612+
pytest.param(
613+
lambda: CustomScheduler(event_function=lambda t: True, start=10, end=50),
614+
id="CustomScheduler",
615+
),
616+
pytest.param(
617+
lambda: MemorylessScheduler(p=0.5, start=10, end=50, seed=42),
618+
id="MemorylessScheduler",
619+
),
620+
pytest.param(
621+
lambda: BurstScheduler(on_duration=3, off_duration=2, start=10, end=50),
622+
id="BurstScheduler",
623+
),
624+
pytest.param(
625+
lambda: DecayingProbabilityScheduler(
626+
initial_probability=0.5, decay_rate=0.01, start=10, end=50, seed=42
627+
),
628+
id="DecayingProbabilityScheduler",
629+
),
630+
pytest.param(
631+
lambda: WindowScheduler(windows=[(10, 50)], start=10, end=50),
632+
id="WindowScheduler",
633+
),
634+
]
635+
636+
637+
class TestSchedulerCommonProperties:
638+
"""Parametrized tests for properties every Scheduler must satisfy."""
639+
640+
@pytest.mark.parametrize("make_sched", _scheduler_factories)
641+
def test_returns_bool_type_in_range(self, make_sched):
642+
"""Result must be a bool when t is within [start, end]."""
643+
sched = make_sched()
644+
result = sched(15)
645+
assert isinstance(result, (bool, np.bool_)), (
646+
f"{sched.__class__.__name__} returned {type(result)}, expected bool"
647+
)
648+
649+
@pytest.mark.parametrize("make_sched", _scheduler_factories)
650+
def test_returns_false_before_start(self, make_sched):
651+
"""Must return False when t < start."""
652+
sched = make_sched()
653+
assert sched(0) is False, (
654+
f"{sched.__class__.__name__} did not return False for t=0 < start=10"
655+
)
656+
657+
@pytest.mark.parametrize("make_sched", _scheduler_factories)
658+
def test_returns_false_after_end(self, make_sched):
659+
"""Must return False when t > end."""
660+
sched = make_sched()
661+
assert sched(100) is False, (
662+
f"{sched.__class__.__name__} did not return False for t=100 > end=50"
663+
)
664+
665+
@pytest.mark.parametrize("make_sched", _scheduler_factories)
666+
def test_result_is_never_none(self, make_sched):
667+
"""__call__ must never return None — always True or False."""
668+
sched = make_sched()
669+
for t in [0, 5, 10, 15, 30, 50, 60]:
670+
result = sched(t)
671+
assert result is not None, (
672+
f"{sched.__class__.__name__} returned None at t={t}"
673+
)
674+
675+
@pytest.mark.parametrize("make_sched", _scheduler_factories)
676+
def test_outside_range_returns_python_bool(self, make_sched):
677+
"""The base class returns literal `False` outside range — must be Python bool."""
678+
sched = make_sched()
679+
result = sched(0)
680+
assert type(result) is bool, (
681+
f"{sched.__class__.__name__} returned {type(result)} outside range, expected bool"
682+
)
683+
684+
@pytest.mark.parametrize("make_sched", _scheduler_factories)
685+
def test_has_start_and_end_attributes(self, make_sched):
686+
"""Every scheduler must expose .start and .end."""
687+
sched = make_sched()
688+
assert hasattr(sched, "start")
689+
assert hasattr(sched, "end")
690+
assert sched.start == 10
691+
assert sched.end == 50
692+
693+
@pytest.mark.parametrize("make_sched", _scheduler_factories)
694+
def test_boundary_start_is_in_range(self, make_sched):
695+
"""t == start should be considered in range (not return False from bounds check)."""
696+
sched = make_sched()
697+
result = sched(10)
698+
# Can't assert True (stochastic schedulers may return False), but must not be None
699+
assert isinstance(result, (bool, np.bool_))
700+
701+
@pytest.mark.parametrize("make_sched", _scheduler_factories)
702+
def test_boundary_end_is_in_range(self, make_sched):
703+
"""t == end should be considered in range (not return False from bounds check)."""
704+
sched = make_sched()
705+
result = sched(50)
706+
assert isinstance(result, (bool, np.bool_))
707+
708+
591709
if __name__ == "__main__":
592710
pytest.main([__file__, "-v"])

tests/test_update_functions.py

Lines changed: 99 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -1014,71 +1014,122 @@ def test_periodic_scheduler_integration(self):
10141014

10151015

10161016
# ============================================================
1017-
# Regression: RandomWalk* must return scalar, not numpy array
1018-
# https://github.com/scope-lab-vu/ns_gym/issues/XX
1019-
# rng.normal(mu, sigma, 1) returns a 1-element array which
1020-
# causes ValueError in downstream dynamics (e.g. Acrobot).
1017+
# Parametrized common-property tests for all single_param UpdateFns
10211018
# ============================================================
10221019

1023-
class TestRandomWalkReturnsScalar:
1024-
"""All RandomWalk* update functions must return a plain Python scalar,
1025-
not a numpy array, so they can be used directly in environment dynamics."""
1026-
1027-
def test_random_walk_returns_scalar(self):
1028-
fn = RandomWalk(_always_scheduler(), mu=0, sigma=1, seed=42)
1029-
param, _, _ = fn(10.0, 0)
1020+
_single_param_factories = [
1021+
pytest.param(lambda s: NoUpdate(s), id="NoUpdate"),
1022+
pytest.param(lambda s: IncrementUpdate(s, k=1.0), id="IncrementUpdate"),
1023+
pytest.param(lambda s: DecrementUpdate(s, k=1.0), id="DecrementUpdate"),
1024+
pytest.param(lambda s: DeterministicTrend(s, slope=1.0), id="DeterministicTrend"),
1025+
pytest.param(lambda s: RandomWalk(s, mu=0, sigma=1, seed=42), id="RandomWalk"),
1026+
pytest.param(
1027+
lambda s: RandomWalkWithDrift(s, alpha=1.0, mu=0, sigma=1, seed=42),
1028+
id="RandomWalkWithDrift",
1029+
),
1030+
pytest.param(
1031+
lambda s: RandomWalkWithDriftAndTrend(s, alpha=1.0, mu=0, sigma=1, slope=0.5, seed=42),
1032+
id="RandomWalkWithDriftAndTrend",
1033+
),
1034+
pytest.param(
1035+
lambda s: StepWiseUpdate(s, param_list=[float(i) for i in range(20)]),
1036+
id="StepWiseUpdate",
1037+
),
1038+
pytest.param(lambda s: OscillatingUpdate(s, delta=1.0), id="OscillatingUpdate"),
1039+
pytest.param(lambda s: ExponentialDecay(s, decay_rate=0.5), id="ExponentialDecay"),
1040+
pytest.param(lambda s: GeometricProgression(s, r=2.0), id="GeometricProgression"),
1041+
pytest.param(
1042+
lambda s: OrnsteinUhlenbeck(s, theta=0.5, mu=10.0, sigma=0.1, seed=42),
1043+
id="OrnsteinUhlenbeck",
1044+
),
1045+
pytest.param(
1046+
lambda s: SigmoidTransition(s, a=1.0, b=10.0, k=1.0, t0=50),
1047+
id="SigmoidTransition",
1048+
),
1049+
pytest.param(
1050+
lambda s: CyclicUpdate(s, value_list=[10.0, 20.0, 30.0]),
1051+
id="CyclicUpdate",
1052+
),
1053+
pytest.param(
1054+
lambda s: BoundedRandomWalk(s, mu=0, sigma=1, lo=-100, hi=100, seed=42),
1055+
id="BoundedRandomWalk",
1056+
),
1057+
pytest.param(lambda s: PolynomialTrend(s, coeffs=[1.0]), id="PolynomialTrend"),
1058+
pytest.param(
1059+
lambda s: LinearInterpolation(s, start_val=0.0, end_val=10.0, T=100),
1060+
id="LinearInterpolation",
1061+
),
1062+
]
1063+
1064+
1065+
class TestSingleParamCommonProperties:
1066+
"""Parametrized tests for properties every single_param UpdateFn must satisfy."""
1067+
1068+
@pytest.mark.parametrize("make_fn", _single_param_factories)
1069+
def test_returns_scalar(self, make_fn):
1070+
"""Output param must be a scalar, not a numpy array."""
1071+
fn = make_fn(_always_scheduler())
1072+
param, _, _ = fn(10.0, 1)
10301073
assert np.isscalar(param), (
1031-
f"RandomWalk returned {type(param)}, expected scalar"
1074+
f"{fn.__class__.__name__} returned {type(param)}, expected scalar"
10321075
)
10331076

1034-
def test_random_walk_with_drift_returns_scalar(self):
1035-
fn = RandomWalkWithDrift(
1036-
_always_scheduler(), alpha=1.0, mu=0, sigma=1, seed=42
1037-
)
1038-
param, _, _ = fn(10.0, 0)
1039-
assert np.isscalar(param), (
1040-
f"RandomWalkWithDrift returned {type(param)}, expected scalar"
1077+
@pytest.mark.parametrize("make_fn", _single_param_factories)
1078+
def test_returns_three_tuple(self, make_fn):
1079+
"""__call__ must return a (param, changed, delta_change) 3-tuple."""
1080+
fn = make_fn(_always_scheduler())
1081+
result = fn(10.0, 1)
1082+
assert len(result) == 3, (
1083+
f"{fn.__class__.__name__} returned {len(result)}-tuple, expected 3"
10411084
)
10421085

1043-
def test_random_walk_with_drift_and_trend_returns_scalar(self):
1044-
fn = RandomWalkWithDriftAndTrend(
1045-
_always_scheduler(), alpha=1.0, mu=0, sigma=1, slope=0.5, seed=42
1046-
)
1047-
param, _, _ = fn(10.0, 1)
1048-
assert np.isscalar(param), (
1049-
f"RandomWalkWithDriftAndTrend returned {type(param)}, expected scalar"
1086+
@pytest.mark.parametrize("make_fn", _single_param_factories)
1087+
def test_changed_flag_is_zero_or_one(self, make_fn):
1088+
"""The changed flag must be 0 or 1."""
1089+
fn = make_fn(_always_scheduler())
1090+
_, changed, _ = fn(10.0, 1)
1091+
assert changed in (0, 1), (
1092+
f"{fn.__class__.__name__} returned changed={changed}, expected 0 or 1"
10501093
)
10511094

1052-
def test_random_walk_scalar_survives_multiple_steps(self):
1053-
"""Ensure the scalar type is preserved across multiple update steps."""
1054-
fn = RandomWalk(_always_scheduler(), mu=0, sigma=1, seed=42)
1055-
param = 10.0
1056-
for t in range(10):
1057-
param, _, _ = fn(param, t)
1058-
assert np.isscalar(param), (
1059-
f"RandomWalk returned {type(param)} at step {t}, expected scalar"
1060-
)
1095+
@pytest.mark.parametrize("make_fn", _single_param_factories)
1096+
def test_delta_change_is_numeric(self, make_fn):
1097+
"""The delta_change element must be a number."""
1098+
fn = make_fn(_always_scheduler())
1099+
_, _, delta = fn(10.0, 1)
1100+
assert isinstance(delta, (int, float, np.integer, np.floating)), (
1101+
f"{fn.__class__.__name__} returned delta type {type(delta)}, expected numeric"
1102+
)
10611103

1062-
def test_random_walk_with_drift_scalar_survives_multiple_steps(self):
1063-
fn = RandomWalkWithDrift(
1064-
_always_scheduler(), alpha=0.5, mu=0, sigma=1, seed=42
1104+
@pytest.mark.parametrize("make_fn", _single_param_factories)
1105+
def test_changed_is_one_when_scheduler_fires(self, make_fn):
1106+
"""When the scheduler fires, changed must be 1."""
1107+
fn = make_fn(_always_scheduler())
1108+
_, changed, _ = fn(10.0, 1)
1109+
assert changed == 1, (
1110+
f"{fn.__class__.__name__} returned changed={changed}, expected 1"
10651111
)
1066-
param = 10.0
1067-
for t in range(10):
1068-
param, _, _ = fn(param, t)
1069-
assert np.isscalar(param), (
1070-
f"RandomWalkWithDrift returned {type(param)} at step {t}, expected scalar"
1071-
)
10721112

1073-
def test_random_walk_with_drift_and_trend_scalar_survives_multiple_steps(self):
1074-
fn = RandomWalkWithDriftAndTrend(
1075-
_always_scheduler(), alpha=0.5, mu=0, sigma=1, slope=0.1, seed=42
1113+
@pytest.mark.parametrize("make_fn", _single_param_factories)
1114+
def test_no_update_when_scheduler_false(self, make_fn):
1115+
"""When the scheduler does not fire: param unchanged, changed=0, delta=0."""
1116+
fn = make_fn(_never_scheduler())
1117+
param, changed, delta = fn(10.0, 0)
1118+
assert param == 10.0, (
1119+
f"{fn.__class__.__name__} changed param to {param} when scheduler was off"
10761120
)
1121+
assert changed == 0
1122+
assert delta == 0.0
1123+
1124+
@pytest.mark.parametrize("make_fn", _single_param_factories)
1125+
def test_scalar_survives_multiple_steps(self, make_fn):
1126+
"""Param must stay scalar across 10 consecutive update steps."""
1127+
fn = make_fn(_always_scheduler())
10771128
param = 10.0
10781129
for t in range(10):
10791130
param, _, _ = fn(param, t)
10801131
assert np.isscalar(param), (
1081-
f"RandomWalkWithDriftAndTrend returned {type(param)} at step {t}, expected scalar"
1132+
f"{fn.__class__.__name__} returned {type(param)} at step {t}, expected scalar"
10821133
)
10831134

10841135

0 commit comments

Comments
 (0)