Repository navigation
Expand file tree
/
Copy pathtest_stopping.py
More file actions
45 lines (34 loc) · 1.54 KB
/
Copy pathtest_stopping.py
File metadata and controls
45 lines (34 loc) · 1.54 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
import pytest
import torch
from thinker_r.stopping import ConvergenceEarlyExit
def test_stops_independently_after_convergence() -> None:
controller = ConvergenceEarlyExit(min_steps=2, max_steps=5, epsilon=0.01, patience=1)
first = controller.update(torch.tensor([[1.0, 0.0], [1.0, 0.0]]))
assert first.reason == ("continue", "continue")
second = controller.update(torch.tensor([[1.0, 0.0], [0.0, 1.0]]))
assert second.should_stop.tolist() == [True, False]
assert second.reason == ("converged", "continue")
def test_inactive_sample_does_not_advance() -> None:
controller = ConvergenceEarlyExit(min_steps=2, max_steps=3, epsilon=0.1, patience=1)
controller.update(torch.ones(2, 4))
decision = controller.update(torch.ones(2, 4), active=torch.tensor([False, True]))
assert decision.latent_steps.tolist() == [1, 2]
assert decision.reason == ("inactive", "converged")
def test_hard_maximum_always_stops() -> None:
controller = ConvergenceEarlyExit(min_steps=2, max_steps=2, epsilon=0.0, patience=2)
controller.update(torch.tensor([[1.0, 0.0]]))
decision = controller.update(torch.tensor([[0.0, 1.0]]))
assert decision.should_stop.item()
assert decision.reason == ("max_steps",)
@pytest.mark.parametrize(
"kwargs",
[
{"min_steps": 0},
{"min_steps": 4, "max_steps": 3},
{"epsilon": -1.0},
{"patience": 0},
],
)
def test_rejects_invalid_configuration(kwargs: dict[str, float]) -> None:
with pytest.raises(ValueError):
ConvergenceEarlyExit(**kwargs)