Skip to content

Commit 4fae5cb

Browse files
Test failed simulator queries are skipped
1 parent 89c35ce commit 4fae5cb

1 file changed

Lines changed: 27 additions & 0 deletions

File tree

tests/learners/test_learners.py

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,12 @@ def _forward(self, x: torch.Tensor) -> torch.Tensor:
1616
return torch.sin(x)
1717

1818

19+
class FailingSin(Simulator):
20+
def _forward(self, x: torch.Tensor) -> None:
21+
del x
22+
return None
23+
24+
1925
def learners(
2026
*, simulator: Simulator, n_initial_samples: int, adaptive_only: bool
2127
) -> Iterable:
@@ -146,6 +152,27 @@ def run_experiment(
146152
return metrics, summary
147153

148154

155+
def test_failed_simulator_query_is_skipped():
156+
simulator = FailingSin(parameters_range={"x": (0, 5.0)}, output_names=["y"])
157+
x_train = torch.tensor([[0.0], [1.0]])
158+
y_train = torch.sin(x_train)
159+
learner = stream.Random(
160+
simulator=simulator,
161+
emulator=GaussianProcessRBF(x_train, y_train, lr=0.001),
162+
x_train=x_train.clone(),
163+
y_train=y_train.clone(),
164+
p_query=1.0,
165+
fit_from_reinitialized=False,
166+
)
167+
168+
learner.fit(torch.tensor([[2.0]]), random_seed=0)
169+
170+
assert torch.equal(learner.x_train, x_train)
171+
assert torch.equal(learner.y_train, y_train)
172+
assert learner.n_queries == 0
173+
assert learner.metrics["n_queries"] == [0]
174+
175+
149176
def test_learners_sin():
150177
metrics, summary = run_experiment(
151178
simulator=Sin(parameters_range={"x": (0, 50.0)}, output_names=["y"]),

0 commit comments

Comments
 (0)